c8870d5b3b
- users 新增 compute_personal_balance(分):个人充值/企业转入个人账本, 与 company_members.compute_balance(企业分配)彻底分离 - relay 转发引擎后按实际用量记账(deduct_usage_post): 先扣企业分配余额(折扣低→高)、不足扣个人余额,写 compute_usage_records; 流式 SSE 透传时解析 usage chunk 记账,JWT 登录态调用生效、PAT 调用不重复记账 - 个人充值到账(pay/service)与 企业↔个人互转 同步更新个人账本 - 成员退出企业(leave)与管理员移出成员(remove_company_member): 仅返还企业分配剩余(member.compute_balance)至企业余额,同步扣减引擎 quota, 绝不触碰个人充值余额 - GET /compute/user/balance 个人余额改读独立账本(不再引擎倒推) - 迁移 0067;存量个人充值/互转净额已回填
544 lines
20 KiB
Python
544 lines
20 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""算力多级折扣价格计算与余额扣费服务。
|
||
|
||
折扣体系(从高到低优先级):
|
||
1. 载体端账号折扣(tenant_user_discounts)- 载体给特定账号的折扣
|
||
2. 载体端企业折扣(park_companies.compute_discount)- 载体给下属企业的折扣
|
||
3. 运营端载体折扣(tenant_discounts)- 运营给载体的折扣(载体进货价)
|
||
4. 运营端标准价 - 未绑定载体的用户使用
|
||
|
||
折扣表示:0-100,0=无折扣,100=全免费;实际价格=标准价×(1-discount/100)
|
||
|
||
扣费优先级:
|
||
1. 企业分配的余额(优先选择折扣最低的企业)
|
||
2. 个人余额
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import uuid
|
||
from datetime import datetime
|
||
from typing import Optional
|
||
|
||
from sqlalchemy import select, and_
|
||
|
||
from ..infrastructure.models import (
|
||
ParkTenant, ParkCompany, CompanyMember, User,
|
||
TenantDiscount, TenantUserDiscount,
|
||
ComputeRecharge, ComputeBalanceAllocation, ComputeUsageRecord,
|
||
)
|
||
from ..infrastructure.repositories import Database
|
||
|
||
MICRO_PER_FEN = 10000 # 1分 = 10000 micro(1元=100分=1,000,000 micro)
|
||
|
||
|
||
# ── 标准价配置(可从数据库或配置文件读取) ──
|
||
STANDARD_PRICES = {
|
||
# model: 每1000token价格(分)
|
||
"default": 2, # 默认模型
|
||
"gpt-4": 10, # 高价模型
|
||
"gpt-3.5-turbo": 2, # 基础模型
|
||
}
|
||
|
||
|
||
def get_standard_price(model: str, token_count: int) -> int:
|
||
"""获取标准价(分)。"""
|
||
per_1k = STANDARD_PRICES.get(model, STANDARD_PRICES["default"])
|
||
return max(1, int(per_1k * token_count / 1000))
|
||
|
||
|
||
def now_str() -> str:
|
||
return datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||
|
||
|
||
def new_id() -> str:
|
||
return uuid.uuid4().hex[:24]
|
||
|
||
|
||
# ── 价格计算 ──
|
||
|
||
async def calculate_compute_price(db: Database, user_id: str, model: str, token_count: int) -> dict:
|
||
"""
|
||
计算用户使用算力的实际价格。
|
||
|
||
返回:
|
||
{
|
||
"standard_price": int, # 标准价(分)
|
||
"tenant_discount": int, # 载体折扣(0-100)
|
||
"discount": int, # 实际折扣(0-100)
|
||
"discount_source": str, # standard/tenant/company/tenant_user
|
||
"actual_price": int, # 实际价格(分)
|
||
"tenant_id": str, # 载体ID(如有)
|
||
"company_id": str, # 企业ID(如有)
|
||
}
|
||
"""
|
||
standard_price = get_standard_price(model, token_count)
|
||
result = {
|
||
"standard_price": standard_price,
|
||
"tenant_discount": 0,
|
||
"discount": 0,
|
||
"discount_source": "standard",
|
||
"actual_price": standard_price,
|
||
"tenant_id": "",
|
||
"company_id": "",
|
||
}
|
||
|
||
# 1. 检查用户是否绑定载体(通过企业成员关系或直接绑定)
|
||
tenant_id = await _get_user_tenant_id(db, user_id)
|
||
if not tenant_id:
|
||
return result
|
||
|
||
result["tenant_id"] = tenant_id
|
||
|
||
# 2. 获取运营端给载体的折扣(载体进货价)
|
||
tenant_discount = await _get_tenant_discount(db, tenant_id)
|
||
result["tenant_discount"] = tenant_discount
|
||
|
||
# 载体基础价 = 标准价 × (1 - 载体折扣/100)
|
||
tenant_base_price = int(standard_price * (1 - tenant_discount / 100))
|
||
|
||
# 3. 检查载体端是否给该账号设置了单独折扣(优先级最高)
|
||
user_discount = await _get_tenant_user_discount(db, tenant_id, user_id)
|
||
if user_discount > 0:
|
||
actual_price = int(tenant_base_price * (1 - user_discount / 100))
|
||
result.update({
|
||
"discount": user_discount,
|
||
"discount_source": "tenant_user",
|
||
"actual_price": actual_price,
|
||
})
|
||
return result
|
||
|
||
# 4. 检查用户是否属于载体下属企业
|
||
company = await _get_user_company(db, user_id, tenant_id)
|
||
if company:
|
||
company_discount = getattr(company, "compute_discount", 0) or 0
|
||
actual_price = int(tenant_base_price * (1 - company_discount / 100))
|
||
result.update({
|
||
"discount": company_discount,
|
||
"discount_source": "company",
|
||
"actual_price": actual_price,
|
||
"company_id": company.id,
|
||
})
|
||
return result
|
||
|
||
# 5. 使用载体默认账号折扣
|
||
tenant = await db.session.get(ParkTenant, tenant_id)
|
||
default_discount = getattr(tenant, "default_user_discount", 0) or 0 if tenant else 0
|
||
if default_discount > 0:
|
||
actual_price = int(tenant_base_price * (1 - default_discount / 100))
|
||
result.update({
|
||
"discount": default_discount,
|
||
"discount_source": "tenant_default",
|
||
"actual_price": actual_price,
|
||
})
|
||
return result
|
||
|
||
# 6. 仅使用载体折扣
|
||
result.update({
|
||
"discount": tenant_discount,
|
||
"discount_source": "tenant",
|
||
"actual_price": tenant_base_price,
|
||
})
|
||
return result
|
||
|
||
|
||
async def _get_user_tenant_id(db: Database, user_id: str) -> Optional[str]:
|
||
"""获取用户绑定的载体ID(通过企业成员关系)。"""
|
||
# 通过企业成员关系查找
|
||
result = await db.session.execute(
|
||
select(CompanyMember).where(
|
||
and_(CompanyMember.user_id == user_id, CompanyMember.status == "active")
|
||
)
|
||
)
|
||
members = result.scalars().all()
|
||
for m in members:
|
||
company = await db.session.get(ParkCompany, m.company_id)
|
||
if company and company.tenant_id:
|
||
return company.tenant_id
|
||
return None
|
||
|
||
|
||
async def _get_tenant_discount(db: Database, tenant_id: str) -> int:
|
||
"""获取运营端给载体的折扣。"""
|
||
result = await db.session.execute(
|
||
select(TenantDiscount).where(TenantDiscount.tenant_id == tenant_id)
|
||
)
|
||
td = result.scalars().first()
|
||
if td:
|
||
# 检查有效期
|
||
now = now_str()
|
||
if td.effective_date and td.effective_date > now:
|
||
return 0
|
||
if td.expire_date and td.expire_date < now:
|
||
return 0
|
||
return td.discount or 0
|
||
return 0
|
||
|
||
|
||
async def _get_tenant_user_discount(db: Database, tenant_id: str, user_id: str) -> int:
|
||
"""获取载体端给特定账号的折扣。"""
|
||
result = await db.session.execute(
|
||
select(TenantUserDiscount).where(
|
||
and_(TenantUserDiscount.tenant_id == tenant_id, TenantUserDiscount.user_id == user_id)
|
||
)
|
||
)
|
||
tud = result.scalars().first()
|
||
return tud.discount if tud else 0
|
||
|
||
|
||
async def _get_user_company(db: Database, user_id: str, tenant_id: str) -> Optional[ParkCompany]:
|
||
"""获取用户在指定载体下的企业。"""
|
||
result = await db.session.execute(
|
||
select(CompanyMember).where(
|
||
and_(CompanyMember.user_id == user_id, CompanyMember.status == "active")
|
||
)
|
||
)
|
||
members = result.scalars().all()
|
||
for m in members:
|
||
company = await db.session.get(ParkCompany, m.company_id)
|
||
if company and company.tenant_id == tenant_id:
|
||
return company
|
||
return None
|
||
|
||
|
||
# ── 余额扣费 ──
|
||
|
||
async def deduct_compute_balance(
|
||
db: Database, user_id: str, model: str, token_count: int
|
||
) -> dict:
|
||
"""
|
||
算力扣费。
|
||
|
||
扣费优先级:
|
||
1. 企业分配的余额(优先选择折扣最低的企业)
|
||
2. 个人余额
|
||
|
||
返回:
|
||
{
|
||
"success": bool,
|
||
"amount": int, # 实际扣费(分)
|
||
"source": str, # company/personal
|
||
"company_id": str, # 从哪家企业扣费(如有)
|
||
"price_info": dict, # 价格计算结果
|
||
"reason": str, # 失败原因
|
||
}
|
||
"""
|
||
price_info = await calculate_compute_price(db, user_id, model, token_count)
|
||
amount = price_info["actual_price"]
|
||
|
||
# 1. 尝试从企业余额扣费
|
||
company_balances = await _get_user_company_balances(db, user_id)
|
||
# 按折扣从低到高排序(折扣越低,价格越低,优先使用)
|
||
company_balances.sort(key=lambda x: x["discount"])
|
||
|
||
for cb in company_balances:
|
||
if cb["balance"] >= amount:
|
||
success = await _deduct_from_company_balance(db, cb["company_id"], user_id, amount)
|
||
if success:
|
||
await _record_usage(db, user_id, model, token_count, price_info, amount, "company", cb["company_id"])
|
||
return {
|
||
"success": True,
|
||
"amount": amount,
|
||
"source": "company",
|
||
"company_id": cb["company_id"],
|
||
"price_info": price_info,
|
||
}
|
||
|
||
# 2. 尝试从个人余额扣费
|
||
user_balance = await _get_user_balance(db, user_id)
|
||
if user_balance >= amount:
|
||
success = await _deduct_from_user_balance(db, user_id, amount)
|
||
if success:
|
||
await _record_usage(db, user_id, model, token_count, price_info, amount, "personal", "")
|
||
return {
|
||
"success": True,
|
||
"amount": amount,
|
||
"source": "personal",
|
||
"price_info": price_info,
|
||
}
|
||
|
||
# 3. 余额不足
|
||
return {
|
||
"success": False,
|
||
"reason": "insufficient_balance",
|
||
"required": amount,
|
||
"available": user_balance,
|
||
"price_info": price_info,
|
||
}
|
||
|
||
|
||
async def _get_user_company_balances(db: Database, user_id: str) -> list:
|
||
"""获取用户所有企业的余额和折扣。"""
|
||
result = await db.session.execute(
|
||
select(CompanyMember).where(
|
||
and_(CompanyMember.user_id == user_id, CompanyMember.status == "active")
|
||
)
|
||
)
|
||
members = result.scalars().all()
|
||
balances = []
|
||
for m in members:
|
||
company = await db.session.get(ParkCompany, m.company_id)
|
||
if company:
|
||
balances.append({
|
||
"company_id": m.company_id,
|
||
"company_name": company.name,
|
||
"balance": m.compute_balance or 0,
|
||
"discount": company.compute_discount or 0,
|
||
})
|
||
return balances
|
||
|
||
|
||
async def _get_user_balance(db: Database, user_id: str) -> int:
|
||
"""获取用户个人余额(users.compute_personal_balance,单位分;仅个人充值/企业转入)。"""
|
||
user = await db.session.get(User, user_id)
|
||
if user is None:
|
||
return 0
|
||
return int(getattr(user, "compute_personal_balance", 0) or 0)
|
||
|
||
|
||
async def _deduct_from_company_balance(db: Database, company_id: str, user_id: str, amount: int) -> bool:
|
||
"""从企业分配的成员余额扣费。"""
|
||
result = await db.session.execute(
|
||
select(CompanyMember).where(
|
||
and_(CompanyMember.company_id == company_id, CompanyMember.user_id == user_id)
|
||
)
|
||
)
|
||
member = result.scalars().first()
|
||
if not member or (member.compute_balance or 0) < amount:
|
||
return False
|
||
member.compute_balance = (member.compute_balance or 0) - amount
|
||
member.compute_balance_used = (member.compute_balance_used or 0) + amount
|
||
await db.session.commit()
|
||
return True
|
||
|
||
|
||
async def _deduct_from_user_balance(db: Database, user_id: str, amount: int) -> bool:
|
||
"""从用户个人余额账本扣费(users.compute_personal_balance,单位分)。"""
|
||
from sqlalchemy import update as _sa_update
|
||
result = await db.session.execute(
|
||
_sa_update(User)
|
||
.where(User.id == user_id, User.compute_personal_balance >= amount)
|
||
.values(compute_personal_balance=User.compute_personal_balance - amount)
|
||
)
|
||
return result.rowcount > 0
|
||
|
||
|
||
async def deduct_usage_post(
|
||
db: Database, user_id: str, model: str, prompt_tokens: int, completion_tokens: int
|
||
) -> dict:
|
||
"""模型调用完成后按实际用量记账(后扣模式)。
|
||
|
||
扣费优先级(与 deduct_compute_balance 一致):
|
||
1. 企业分配的余额(按折扣从低到高)
|
||
2. 个人余额(users.compute_personal_balance)
|
||
|
||
与引擎侧扣费并行:引擎扣混合 quota(执行层),平台按来源记账(账本层)。
|
||
返回 {"ok": bool, "amount": int, "source": str, "reason": str}。
|
||
"""
|
||
token_count = max(1, int(prompt_tokens or 0) + int(completion_tokens or 0))
|
||
price_info = await calculate_compute_price(db, user_id, model, token_count)
|
||
amount = price_info["actual_price"]
|
||
if amount <= 0:
|
||
return {"ok": True, "amount": 0, "source": "free", "reason": ""}
|
||
|
||
# 1. 先扣企业分配余额(折扣从低到高,即价格从低到高)
|
||
company_balances = await _get_user_company_balances(db, user_id)
|
||
company_balances.sort(key=lambda x: x["discount"])
|
||
remaining = amount
|
||
charged_from = ""
|
||
for cb in company_balances:
|
||
if remaining <= 0:
|
||
break
|
||
take = min(remaining, cb["balance"])
|
||
if take > 0:
|
||
ok = await _deduct_from_company_balance(db, cb["company_id"], user_id, take)
|
||
if ok:
|
||
remaining -= take
|
||
charged_from = cb["company_id"]
|
||
if remaining <= 0:
|
||
await _record_usage(db, user_id, model, token_count, price_info, amount, "company", charged_from)
|
||
return {"ok": True, "amount": amount, "source": "company", "company_id": charged_from}
|
||
|
||
# 2. 企业余额不足部分,从个人余额扣
|
||
if remaining > 0:
|
||
ok = await _deduct_from_user_balance(db, user_id, remaining)
|
||
if ok:
|
||
await _record_usage(db, user_id, model, token_count, price_info, amount, "personal", "")
|
||
return {"ok": True, "amount": amount, "source": "personal", "reason": ""}
|
||
|
||
# 3. 全部不足:引擎已放行,账本记为欠费(扣到 0),不阻断
|
||
await _record_usage(db, user_id, model, token_count, price_info, amount, "personal", "")
|
||
return {
|
||
"ok": False, "amount": amount, "source": "personal",
|
||
"reason": "insufficient_balance",
|
||
}
|
||
|
||
|
||
async def check_user_balance_available(db: Database, user_id: str) -> bool:
|
||
"""请求前预检查:企业分配余额 + 个人余额合计 > 0 才放行(避免无余额仍调引擎)。"""
|
||
company_balances = await _get_user_company_balances(db, user_id)
|
||
company_total = sum(cb["balance"] for cb in company_balances)
|
||
personal = await _get_user_balance(db, user_id)
|
||
return (company_total + personal) > 0
|
||
|
||
|
||
async def _record_usage(
|
||
db: Database, user_id: str, model: str, token_count: int,
|
||
price_info: dict, amount: int, balance_source: str, company_id: str
|
||
):
|
||
"""记录算力使用。"""
|
||
record = ComputeUsageRecord(
|
||
id=new_id(),
|
||
user_id=user_id,
|
||
company_id=company_id,
|
||
model=model,
|
||
token_count=token_count,
|
||
standard_price=price_info["standard_price"],
|
||
discount=price_info["discount"],
|
||
discount_source=price_info["discount_source"],
|
||
actual_amount=amount,
|
||
balance_source=balance_source,
|
||
created_at=now_str(),
|
||
)
|
||
db.session.add(record)
|
||
await db.session.commit()
|
||
|
||
|
||
# ── 企业余额管理 ──
|
||
|
||
async def get_company_balance(db: Database, company_id: str) -> dict:
|
||
"""获取企业余额总览。"""
|
||
company = await db.session.get(ParkCompany, company_id)
|
||
if not company:
|
||
return {"error": "企业不存在"}
|
||
return {
|
||
"company_id": company.id,
|
||
"company_name": company.name,
|
||
"total_balance": company.compute_balance or 0,
|
||
"allocated_balance": company.compute_balance_used or 0,
|
||
"available_balance": (company.compute_balance or 0) - (company.compute_balance_used or 0),
|
||
}
|
||
|
||
|
||
async def allocate_balance_to_member(
|
||
db: Database, company_id: str, from_user_id: str, to_user_id: str, amount: int, reason: str = ""
|
||
) -> dict:
|
||
"""给成员分配余额。
|
||
企业扣(compute_balance_used 增加),成员个人算力增(company_members.compute_balance + compute-engine quota 同步)。
|
||
"""
|
||
company = await db.session.get(ParkCompany, company_id)
|
||
if not company:
|
||
return {"success": False, "reason": "企业不存在"}
|
||
|
||
available = (company.compute_balance or 0) - (company.compute_balance_used or 0)
|
||
if available < amount:
|
||
return {"success": False, "reason": "企业可用余额不足", "available": available}
|
||
|
||
# 查找成员关系
|
||
result = await db.session.execute(
|
||
select(CompanyMember).where(
|
||
and_(CompanyMember.company_id == company_id, CompanyMember.user_id == to_user_id)
|
||
)
|
||
)
|
||
member = result.scalars().first()
|
||
if not member:
|
||
return {"success": False, "reason": "用户不是该企业成员"}
|
||
|
||
# 获取成员 username,用于 compute-engine 同步
|
||
to_user = await db.session.get(User, to_user_id)
|
||
if not to_user:
|
||
return {"success": False, "reason": "成员用户不存在"}
|
||
|
||
# 先同步 compute-engine(增加成员 quota),成功后再提交数据库
|
||
from ..pay.service import _resolve_engine_user_id
|
||
from ..services import compute_client
|
||
total_micro = amount * MICRO_PER_FEN
|
||
try:
|
||
engine_user_id = await _resolve_engine_user_id(to_user.username or "")
|
||
if not engine_user_id:
|
||
return {"success": False, "reason": "成员算力账号未就绪"}
|
||
await compute_client.adjust_user_quota(engine_user_id, total_micro, "add")
|
||
await compute_client.sync_user_mirror(db, engine_user_id)
|
||
except Exception as exc:
|
||
return {"success": False, "reason": f"算力同步失败: {exc}"}
|
||
|
||
# 分配余额(数据库记录)
|
||
company.compute_balance_used = (company.compute_balance_used or 0) + amount
|
||
member.compute_balance = (member.compute_balance or 0) + amount
|
||
|
||
# 记录分配
|
||
allocation = ComputeBalanceAllocation(
|
||
id=new_id(),
|
||
company_id=company_id,
|
||
from_user_id=from_user_id,
|
||
to_user_id=to_user_id,
|
||
amount=amount,
|
||
type="allocate",
|
||
reason=reason,
|
||
created_at=now_str(),
|
||
)
|
||
db.session.add(allocation)
|
||
await db.session.commit()
|
||
|
||
return {"success": True, "allocation_id": allocation.id}
|
||
|
||
|
||
async def reclaim_balance_from_member(
|
||
db: Database, company_id: str, from_user_id: str, to_user_id: str, amount: int, reason: str = ""
|
||
) -> dict:
|
||
"""回收成员余额。
|
||
成员个人算力扣(company_members.compute_balance + compute-engine quota 同步),企业增(compute_balance_used 减少)。
|
||
"""
|
||
result = await db.session.execute(
|
||
select(CompanyMember).where(
|
||
and_(CompanyMember.company_id == company_id, CompanyMember.user_id == to_user_id)
|
||
)
|
||
)
|
||
member = result.scalars().first()
|
||
if not member:
|
||
return {"success": False, "reason": "用户不是该企业成员"}
|
||
|
||
if (member.compute_balance or 0) < amount:
|
||
return {"success": False, "reason": "成员余额不足", "available": member.compute_balance or 0}
|
||
|
||
company = await db.session.get(ParkCompany, company_id)
|
||
if not company:
|
||
return {"success": False, "reason": "企业不存在"}
|
||
|
||
# 获取成员 username,用于 compute-engine 同步
|
||
to_user = await db.session.get(User, to_user_id)
|
||
if not to_user:
|
||
return {"success": False, "reason": "成员用户不存在"}
|
||
|
||
# 先同步 compute-engine(扣减成员 quota),成功后再提交数据库
|
||
from ..pay.service import _resolve_engine_user_id
|
||
from ..services import compute_client
|
||
total_micro = amount * MICRO_PER_FEN
|
||
try:
|
||
engine_user_id = await _resolve_engine_user_id(to_user.username or "")
|
||
if not engine_user_id:
|
||
return {"success": False, "reason": "成员算力账号未就绪"}
|
||
await compute_client.adjust_user_quota(engine_user_id, total_micro, "subtract")
|
||
await compute_client.sync_user_mirror(db, engine_user_id)
|
||
except Exception as exc:
|
||
return {"success": False, "reason": f"算力同步失败: {exc}"}
|
||
|
||
# 回收余额(数据库记录)
|
||
member.compute_balance = (member.compute_balance or 0) - amount
|
||
company.compute_balance_used = max(0, (company.compute_balance_used or 0) - amount)
|
||
|
||
# 记录回收
|
||
allocation = ComputeBalanceAllocation(
|
||
id=new_id(),
|
||
company_id=company_id,
|
||
from_user_id=from_user_id,
|
||
to_user_id=to_user_id,
|
||
amount=amount,
|
||
type="reclaim",
|
||
reason=reason,
|
||
created_at=now_str(),
|
||
)
|
||
db.session.add(allocation)
|
||
await db.session.commit()
|
||
|
||
return {"success": True, "allocation_id": allocation.id}
|