1076 lines
43 KiB
Python
1076 lines
43 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""算力多级折扣与余额管理 API。
|
||
|
||
包含:
|
||
- 运营端:载体折扣管理、标准价管理
|
||
- 载体端:企业折扣设置、账号折扣设置
|
||
- 企业端:余额充值、成员余额分配
|
||
- 个人端:我的余额、退出企业
|
||
- 内部:价格查询、算力扣费
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
import uuid
|
||
from datetime import datetime
|
||
from typing import Optional
|
||
|
||
import httpx
|
||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||
from pydantic import BaseModel
|
||
from sqlalchemy import select, and_
|
||
|
||
from ..dependencies import get_db, get_current_user
|
||
from ...rbac import write_audit
|
||
from ...domain.rules import role_allowed
|
||
from ...config import COMPUTE_BASE_URL, COMPUTE_ADMIN_TOKEN, COMPUTE_TIMEOUT
|
||
from ...infrastructure.models import (
|
||
ParkTenant, ParkCompany, CompanyMember, User,
|
||
TenantDiscount, TenantUserDiscount,
|
||
ComputeRecharge, ComputeBalanceAllocation, ComputeUsageRecord,
|
||
)
|
||
from ...pay.models import ComputeRechargeOrder
|
||
from ...infrastructure.repositories import Database
|
||
from ...services.compute_pricing_service import (
|
||
calculate_compute_price, deduct_compute_balance,
|
||
get_company_balance, allocate_balance_to_member, reclaim_balance_from_member,
|
||
STANDARD_PRICES,
|
||
)
|
||
|
||
router = APIRouter(prefix="/compute", tags=["compute-pricing"])
|
||
logger = logging.getLogger("compute-pricing")
|
||
|
||
|
||
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 _write_transfer_bill(
|
||
db: Database,
|
||
*,
|
||
target_type: str, # "user" | "company"
|
||
target_id: str,
|
||
amount_fen: int,
|
||
source_label: str, # 来源标注:企业转入 / 个人转入 / 企业分配 / 企业回收
|
||
operator_id: str = "",
|
||
from_name: str = "",
|
||
):
|
||
"""互转/分配时写入账单记录,确保接收方的充值账单中有记录。
|
||
- target_type=user: 写入 compute_recharge_orders(个人账单)
|
||
- target_type=company: 写入 compute_recharges(企业账单)
|
||
"""
|
||
now = now_str()
|
||
quota_micro = amount_fen * MICRO_PER_FEN # 分→micro
|
||
|
||
if target_type == "user":
|
||
# 个人账单:写入 compute_recharge_orders
|
||
order = ComputeRechargeOrder(
|
||
id=f"tr_{uuid.uuid4().hex[:16]}",
|
||
order_no=f"TR_{int(datetime.now().timestamp())}_{uuid.uuid4().hex[:8]}",
|
||
user_id=target_id,
|
||
target_type="user",
|
||
target_id=target_id,
|
||
package_id=source_label,
|
||
amount_fen=amount_fen,
|
||
quota_micro=quota_micro,
|
||
status="credited",
|
||
created_at=now,
|
||
paid_at=now,
|
||
credited_at=now,
|
||
notify_payload=f'{{"source":"{source_label}","from":"{from_name}","operator":"{operator_id}"}}',
|
||
)
|
||
db.session.add(order)
|
||
elif target_type == "company":
|
||
# 企业账单:写入 compute_recharges
|
||
recharge = ComputeRecharge(
|
||
id=new_id(),
|
||
company_id=target_id,
|
||
user_id=operator_id,
|
||
amount=amount_fen,
|
||
compute_amount=amount_fen,
|
||
payment_method=source_label,
|
||
payment_status="paid",
|
||
recharged_by=operator_id,
|
||
created_at=now,
|
||
paid_at=now,
|
||
)
|
||
db.session.add(recharge)
|
||
|
||
|
||
def check_role(user: dict, roles: list[str]) -> None:
|
||
"""检查用户角色,如果不允许则抛出403异常。"""
|
||
if not role_allowed(user, *roles):
|
||
raise HTTPException(status_code=403, detail="Forbidden: insufficient role")
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
# 运营端:载体折扣管理
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
|
||
class TenantDiscountBody(BaseModel):
|
||
tenant_id: str
|
||
discount: int = 0
|
||
effective_date: str = ""
|
||
expire_date: str = ""
|
||
reason: str = ""
|
||
|
||
|
||
@router.get("/admin/tenant-discounts", summary="运营端:载体折扣列表")
|
||
async def admin_tenant_discounts(
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
check_role(user, ["operator"])
|
||
result = await db.session.execute(select(TenantDiscount).order_by(TenantDiscount.created_at.desc()))
|
||
items = []
|
||
for td in result.scalars().all():
|
||
tenant = await db.session.get(ParkTenant, td.tenant_id)
|
||
items.append({
|
||
"id": td.id,
|
||
"tenant_id": td.tenant_id,
|
||
"tenant_name": tenant.name if tenant else "",
|
||
"discount": td.discount,
|
||
"effective_date": td.effective_date,
|
||
"expire_date": td.expire_date,
|
||
"reason": td.reason,
|
||
"created_by": td.created_by,
|
||
"created_at": td.created_at,
|
||
})
|
||
return {"items": items}
|
||
|
||
|
||
@router.post("/admin/tenant-discounts", summary="运营端:设置载体折扣")
|
||
async def admin_set_tenant_discount(
|
||
body: TenantDiscountBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
check_role(user, ["operator"])
|
||
body.discount = max(0, min(100, int(body.discount)))
|
||
|
||
# 检查载体是否存在
|
||
tenant = await db.session.get(ParkTenant, body.tenant_id)
|
||
if not tenant:
|
||
raise HTTPException(status_code=404, detail="载体不存在")
|
||
|
||
# 检查是否已有折扣,有则更新
|
||
result = await db.session.execute(
|
||
select(TenantDiscount).where(TenantDiscount.tenant_id == body.tenant_id)
|
||
)
|
||
td = result.scalars().first()
|
||
if td:
|
||
td.discount = body.discount
|
||
td.effective_date = body.effective_date
|
||
td.expire_date = body.expire_date
|
||
td.reason = body.reason
|
||
td.created_by = user.get("id", "")
|
||
td.created_at = now_str()
|
||
else:
|
||
td = TenantDiscount(
|
||
id=new_id(),
|
||
tenant_id=body.tenant_id,
|
||
discount=body.discount,
|
||
effective_date=body.effective_date,
|
||
expire_date=body.expire_date,
|
||
reason=body.reason,
|
||
created_by=user.get("id", ""),
|
||
created_at=now_str(),
|
||
)
|
||
db.session.add(td)
|
||
|
||
await db.session.commit()
|
||
await write_audit(db, action="compute.tenant_discount_set", resource="tenant_discount",
|
||
resource_id=td.id, detail=f"tenant={body.tenant_id}, discount={body.discount}",
|
||
user=user, request=request)
|
||
return {"ok": True, "id": td.id}
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
# 运营端:标准价管理
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
|
||
class StandardPriceBody(BaseModel):
|
||
model: str
|
||
price_per_1k: int # 每1000token价格(分)
|
||
|
||
|
||
@router.get("/admin/standard-prices", summary="运营端:标准价列表")
|
||
async def admin_standard_prices(
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
check_role(user, ["operator"])
|
||
items = [{"model": k, "price_per_1k": v} for k, v in STANDARD_PRICES.items()]
|
||
return {"items": items}
|
||
|
||
|
||
@router.put("/admin/standard-prices/{model}", summary="运营端:修改标准价")
|
||
async def admin_update_standard_price(
|
||
model: str,
|
||
body: StandardPriceBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
check_role(user, ["operator"])
|
||
new_price = max(1, int(body.price_per_1k))
|
||
STANDARD_PRICES[model] = new_price
|
||
|
||
# 同步更新 compute 服务 models 表价格(实际计费以此为准)
|
||
# 单位转换:分/1000token → 元/百万token = price * 10
|
||
compute_price_yuan_per_m = round(new_price * 10, 4)
|
||
synced = await _sync_compute_model_price(model, compute_price_yuan_per_m)
|
||
|
||
await write_audit(db, action="compute.standard_price_update", resource="standard_price",
|
||
resource_id=model, detail=f"price={new_price}, compute_synced={synced}",
|
||
user=user, request=request)
|
||
return {"ok": True, "model": model, "price_per_1k": STANDARD_PRICES[model],
|
||
"compute_synced": synced}
|
||
|
||
|
||
async def _sync_compute_model_price(model_name: str, price_yuan_per_m: float) -> bool:
|
||
"""同步价格到 compute 服务 models 表(input_price = output_price = 统一价)。
|
||
|
||
返回是否同步成功;失败不影响 server-core 内存价格更新,仅记录日志。
|
||
"""
|
||
if not COMPUTE_BASE_URL or not COMPUTE_ADMIN_TOKEN:
|
||
logger.warning("compute 服务地址或 token 未配置,跳过价格同步")
|
||
return False
|
||
headers = {"Authorization": f"Bearer {COMPUTE_ADMIN_TOKEN}",
|
||
"Content-Type": "application/json"}
|
||
try:
|
||
async with httpx.AsyncClient(timeout=COMPUTE_TIMEOUT) as client:
|
||
# 1. 查模型列表,按 model_name 匹配 id
|
||
resp = await client.get(f"{COMPUTE_BASE_URL}/admin/models",
|
||
params={"page_size": 200}, headers=headers)
|
||
if resp.status_code != 200:
|
||
logger.warning(f"compute 模型列表查询失败: {resp.status_code}")
|
||
return False
|
||
data = resp.json()
|
||
items = data.get("data", {}).get("items", data.get("data", []))
|
||
if isinstance(items, dict):
|
||
items = items.get("items", [])
|
||
target = None
|
||
for m in items:
|
||
if m.get("model_name") == model_name or m.get("name") == model_name:
|
||
target = m
|
||
break
|
||
if target is None:
|
||
logger.info(f"compute 服务中未找到模型 {model_name},跳过同步")
|
||
return False
|
||
# 2. 更新价格
|
||
mid = target.get("id")
|
||
resp = await client.put(f"{COMPUTE_BASE_URL}/admin/models",
|
||
json={"id": mid,
|
||
"input_price": price_yuan_per_m,
|
||
"output_price": price_yuan_per_m},
|
||
headers=headers)
|
||
if resp.status_code == 200:
|
||
logger.info(f"compute 模型 {model_name}(id={mid}) 价格已同步为 {price_yuan_per_m} 元/百万token")
|
||
return True
|
||
logger.warning(f"compute 模型价格更新失败: {resp.status_code} {resp.text}")
|
||
return False
|
||
except Exception as e:
|
||
logger.warning(f"compute 价格同步异常: {e}")
|
||
return False
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
# 载体端:折扣设置
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
|
||
class TenantDiscountSettingsBody(BaseModel):
|
||
default_company_discount: int = 0
|
||
default_user_discount: int = 0
|
||
|
||
|
||
class CompanyDiscountBody(BaseModel):
|
||
discount: int = 0
|
||
|
||
|
||
class UserDiscountBody(BaseModel):
|
||
user_id: str
|
||
discount: int = 0
|
||
reason: str = ""
|
||
|
||
|
||
async def _get_current_tenant(db: Database, user: dict) -> Optional[ParkTenant]:
|
||
"""获取当前用户管理的载体。"""
|
||
# 通过operator_user_id查找
|
||
result = await db.session.execute(
|
||
select(ParkTenant).where(ParkTenant.operator_user_id == user.get("id", ""))
|
||
)
|
||
return result.scalars().first()
|
||
|
||
|
||
@router.get("/tenant/discount-settings", summary="载体端:获取折扣设置")
|
||
async def tenant_discount_settings(
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
tenant = await _get_current_tenant(db, user)
|
||
if not tenant:
|
||
raise HTTPException(status_code=403, detail="未绑定载体")
|
||
return {
|
||
"tenant_id": tenant.id,
|
||
"tenant_name": tenant.name,
|
||
"default_company_discount": tenant.default_company_discount or 0,
|
||
"default_user_discount": tenant.default_user_discount or 0,
|
||
}
|
||
|
||
|
||
@router.put("/tenant/discount-settings", summary="载体端:更新默认折扣")
|
||
async def tenant_update_discount_settings(
|
||
body: TenantDiscountSettingsBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
tenant = await _get_current_tenant(db, user)
|
||
if not tenant:
|
||
raise HTTPException(status_code=403, detail="未绑定载体")
|
||
|
||
body.default_company_discount = max(0, min(100, int(body.default_company_discount)))
|
||
body.default_user_discount = max(0, min(100, int(body.default_user_discount)))
|
||
|
||
tenant.default_company_discount = body.default_company_discount
|
||
tenant.default_user_discount = body.default_user_discount
|
||
await db.session.commit()
|
||
|
||
await write_audit(db, action="compute.tenant_discount_settings_update", resource="park_tenant",
|
||
resource_id=tenant.id, user=user, request=request)
|
||
return {"ok": True}
|
||
|
||
|
||
@router.get("/tenant/company-discounts", summary="载体端:下属企业折扣列表")
|
||
async def tenant_company_discounts(
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
tenant = await _get_current_tenant(db, user)
|
||
if not tenant:
|
||
raise HTTPException(status_code=403, detail="未绑定载体")
|
||
|
||
result = await db.session.execute(
|
||
select(ParkCompany).where(ParkCompany.tenant_id == tenant.id)
|
||
)
|
||
items = []
|
||
for c in result.scalars().all():
|
||
items.append({
|
||
"company_id": c.id,
|
||
"company_name": c.name,
|
||
"discount": c.compute_discount or 0,
|
||
"industry": c.industry,
|
||
"status": c.status,
|
||
})
|
||
return {"items": items}
|
||
|
||
|
||
@router.put("/tenant/company-discounts/{company_id}", summary="载体端:设置企业折扣")
|
||
async def tenant_set_company_discount(
|
||
company_id: str,
|
||
body: CompanyDiscountBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
tenant = await _get_current_tenant(db, user)
|
||
if not tenant:
|
||
raise HTTPException(status_code=403, detail="未绑定载体")
|
||
|
||
company = await db.session.get(ParkCompany, company_id)
|
||
if not company or company.tenant_id != tenant.id:
|
||
raise HTTPException(status_code=404, detail="企业不存在或不属于本载体")
|
||
|
||
body.discount = max(0, min(100, int(body.discount)))
|
||
company.compute_discount = body.discount
|
||
await db.session.commit()
|
||
|
||
await write_audit(db, action="compute.company_discount_set", resource="park_company",
|
||
resource_id=company_id, detail=f"discount={body.discount}",
|
||
user=user, request=request)
|
||
return {"ok": True}
|
||
|
||
|
||
@router.get("/tenant/user-discounts", summary="载体端:下属账号折扣列表")
|
||
async def tenant_user_discounts(
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
tenant = await _get_current_tenant(db, user)
|
||
if not tenant:
|
||
raise HTTPException(status_code=403, detail="未绑定载体")
|
||
|
||
result = await db.session.execute(
|
||
select(TenantUserDiscount).where(TenantUserDiscount.tenant_id == tenant.id)
|
||
)
|
||
items = []
|
||
for tud in result.scalars().all():
|
||
items.append({
|
||
"id": tud.id,
|
||
"user_id": tud.user_id,
|
||
"discount": tud.discount,
|
||
"reason": tud.reason,
|
||
"created_at": tud.created_at,
|
||
})
|
||
return {"items": items}
|
||
|
||
|
||
@router.post("/tenant/user-discounts", summary="载体端:设置账号折扣")
|
||
async def tenant_set_user_discount(
|
||
body: UserDiscountBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
tenant = await _get_current_tenant(db, user)
|
||
if not tenant:
|
||
raise HTTPException(status_code=403, detail="未绑定载体")
|
||
|
||
body.discount = max(0, min(100, int(body.discount)))
|
||
|
||
# 检查是否已有折扣
|
||
result = await db.session.execute(
|
||
select(TenantUserDiscount).where(
|
||
and_(TenantUserDiscount.tenant_id == tenant.id, TenantUserDiscount.user_id == body.user_id)
|
||
)
|
||
)
|
||
tud = result.scalars().first()
|
||
if tud:
|
||
tud.discount = body.discount
|
||
tud.reason = body.reason
|
||
tud.created_by = user.get("id", "")
|
||
tud.created_at = now_str()
|
||
else:
|
||
tud = TenantUserDiscount(
|
||
id=new_id(),
|
||
tenant_id=tenant.id,
|
||
user_id=body.user_id,
|
||
discount=body.discount,
|
||
reason=body.reason,
|
||
created_by=user.get("id", ""),
|
||
created_at=now_str(),
|
||
)
|
||
db.session.add(tud)
|
||
|
||
await db.session.commit()
|
||
await write_audit(db, action="compute.user_discount_set", resource="tenant_user_discount",
|
||
resource_id=tud.id, detail=f"user={body.user_id}, discount={body.discount}",
|
||
user=user, request=request)
|
||
return {"ok": True, "id": tud.id}
|
||
|
||
|
||
@router.delete("/tenant/user-discounts/{discount_id}", summary="载体端:删除账号折扣")
|
||
async def tenant_delete_user_discount(
|
||
discount_id: str,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
tenant = await _get_current_tenant(db, user)
|
||
if not tenant:
|
||
raise HTTPException(status_code=403, detail="未绑定载体")
|
||
|
||
tud = await db.session.get(TenantUserDiscount, discount_id)
|
||
if not tud or tud.tenant_id != tenant.id:
|
||
raise HTTPException(status_code=404, detail="折扣记录不存在")
|
||
|
||
await db.session.delete(tud)
|
||
await db.session.commit()
|
||
await write_audit(db, action="compute.user_discount_delete", resource="tenant_user_discount",
|
||
resource_id=discount_id, user=user, request=request)
|
||
return {"ok": True}
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
# 企业端:余额管理
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
|
||
class RechargeBody(BaseModel):
|
||
amount: int # 充值金额(分)
|
||
payment_method: str = "wechat"
|
||
|
||
|
||
class AllocateBody(BaseModel):
|
||
user_id: str
|
||
amount: int
|
||
reason: str = ""
|
||
|
||
|
||
async def _get_admin_company(db: Database, company_id: str, user: dict) -> ParkCompany:
|
||
"""校验当前用户是该企业的管理员。"""
|
||
result = await db.session.execute(
|
||
select(CompanyMember).where(
|
||
and_(CompanyMember.company_id == company_id, CompanyMember.user_id == user.get("id", ""))
|
||
)
|
||
)
|
||
member = result.scalars().first()
|
||
if not member or not member.is_admin:
|
||
raise HTTPException(status_code=403, detail="仅企业管理员可操作")
|
||
|
||
company = await db.session.get(ParkCompany, company_id)
|
||
if not company:
|
||
raise HTTPException(status_code=404, detail="企业不存在")
|
||
return company
|
||
|
||
|
||
@router.get("/enterprise/{company_id}/balance", summary="企业端:余额总览")
|
||
async def enterprise_balance(
|
||
company_id: str,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
await _get_admin_company(db, company_id, user)
|
||
return await get_company_balance(db, company_id)
|
||
|
||
|
||
@router.post("/enterprise/{company_id}/recharge", summary="企业端:充值")
|
||
async def enterprise_recharge(
|
||
company_id: str,
|
||
body: RechargeBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
company = await _get_admin_company(db, company_id, user)
|
||
|
||
if body.amount <= 0:
|
||
raise HTTPException(status_code=400, detail="充值金额必须大于0")
|
||
|
||
# 创建充值记录(模拟支付成功)
|
||
recharge = ComputeRecharge(
|
||
id=new_id(),
|
||
company_id=company_id,
|
||
amount=body.amount,
|
||
compute_amount=body.amount, # 1:1到账
|
||
payment_method=body.payment_method,
|
||
payment_status="paid",
|
||
recharged_by=user.get("id", ""),
|
||
created_at=now_str(),
|
||
paid_at=now_str(),
|
||
)
|
||
db.session.add(recharge)
|
||
|
||
# 增加企业余额
|
||
company.compute_balance = (company.compute_balance or 0) + body.amount
|
||
await db.session.commit()
|
||
|
||
await write_audit(db, action="compute.enterprise_recharge", resource="compute_recharge",
|
||
resource_id=recharge.id, detail=f"amount={body.amount}",
|
||
user=user, request=request)
|
||
return {"ok": True, "recharge_id": recharge.id, "balance": company.compute_balance}
|
||
|
||
|
||
@router.get("/enterprise/{company_id}/recharges", summary="企业端:充值记录")
|
||
async def enterprise_recharges(
|
||
company_id: str,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
await _get_admin_company(db, company_id, user)
|
||
|
||
# 自动关闭超过 24 小时未支付的 pending 订单
|
||
from datetime import datetime, timedelta
|
||
cutoff = (datetime.now() - timedelta(hours=24)).strftime("%Y-%m-%d %H:%M:%S")
|
||
expired = (await db.session.scalars(
|
||
select(ComputeRecharge).where(
|
||
ComputeRecharge.company_id == company_id,
|
||
ComputeRecharge.payment_status == "pending",
|
||
ComputeRecharge.created_at < cutoff,
|
||
)
|
||
)).all()
|
||
for r in expired:
|
||
r.payment_status = "closed"
|
||
if expired:
|
||
await db.session.commit()
|
||
|
||
# 只返回非 pending 状态的记录(待支付不显示在账单中)
|
||
result = await db.session.execute(
|
||
select(ComputeRecharge).where(
|
||
ComputeRecharge.company_id == company_id,
|
||
ComputeRecharge.payment_status != "pending",
|
||
).order_by(ComputeRecharge.created_at.desc())
|
||
)
|
||
items = [{"id": r.id, "amount": r.amount, "compute_amount": r.compute_amount,
|
||
"payment_status": r.payment_status, "created_at": r.created_at,
|
||
"payment_method": r.payment_method}
|
||
for r in result.scalars().all()]
|
||
return {"items": items}
|
||
|
||
|
||
@router.post("/enterprise/{company_id}/allocate", summary="企业端:给成员分配余额")
|
||
async def enterprise_allocate(
|
||
company_id: str,
|
||
body: AllocateBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
await _get_admin_company(db, company_id, user)
|
||
company = await db.session.get(ParkCompany, company_id)
|
||
result = await allocate_balance_to_member(db, company_id, user.get("id", ""), body.user_id, body.amount, body.reason)
|
||
if not result.get("success"):
|
||
raise HTTPException(status_code=400, detail=result.get("reason", "分配失败"))
|
||
|
||
# 写入成员个人账单(企业分配)
|
||
await _write_transfer_bill(
|
||
db, target_type="user", target_id=body.user_id, amount_fen=body.amount,
|
||
source_label="企业分配", operator_id=user.get("id", ""), from_name=company.name if company else "",
|
||
)
|
||
await db.session.commit()
|
||
|
||
await write_audit(db, action="compute.balance_allocate", resource="compute_balance_allocation",
|
||
resource_id=result["allocation_id"], detail=f"user={body.user_id}, amount={body.amount}",
|
||
user=user, request=request)
|
||
return result
|
||
|
||
|
||
@router.post("/enterprise/{company_id}/reclaim", summary="企业端:回收成员余额")
|
||
async def enterprise_reclaim(
|
||
company_id: str,
|
||
body: AllocateBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
await _get_admin_company(db, company_id, user)
|
||
result = await reclaim_balance_from_member(db, company_id, user.get("id", ""), body.user_id, body.amount, body.reason)
|
||
if not result.get("success"):
|
||
raise HTTPException(status_code=400, detail=result.get("reason", "回收失败"))
|
||
await write_audit(db, action="compute.balance_reclaim", resource="compute_balance_allocation",
|
||
resource_id=result["allocation_id"], detail=f"user={body.user_id}, amount={body.amount}",
|
||
user=user, request=request)
|
||
return result
|
||
|
||
|
||
@router.get("/enterprise/{company_id}/allocations", summary="企业端:分配记录")
|
||
async def enterprise_allocations(
|
||
company_id: str,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
await _get_admin_company(db, company_id, user)
|
||
result = await db.session.execute(
|
||
select(ComputeBalanceAllocation).where(ComputeBalanceAllocation.company_id == company_id).order_by(ComputeBalanceAllocation.created_at.desc())
|
||
)
|
||
items = [{"id": a.id, "from_user_id": a.from_user_id, "to_user_id": a.to_user_id,
|
||
"amount": a.amount, "type": a.type, "reason": a.reason, "created_at": a.created_at}
|
||
for a in result.scalars().all()]
|
||
return {"items": items}
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
# 个人端:我的余额、退出企业
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
|
||
@router.get("/user/balance", summary="个人端:我的余额")
|
||
async def user_balance(
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
user_id = user.get("id", "")
|
||
username = user.get("username", "")
|
||
|
||
# 获取企业分配的余额(company_members.compute_balance 单位=分)
|
||
result = await db.session.execute(
|
||
select(CompanyMember).where(
|
||
and_(CompanyMember.user_id == user_id, CompanyMember.status == "active")
|
||
)
|
||
)
|
||
company_balances = []
|
||
total_company_balance_fen = 0
|
||
for m in result.scalars().all():
|
||
company = await db.session.get(ParkCompany, m.company_id)
|
||
if company:
|
||
balance_fen = m.compute_balance or 0
|
||
total_company_balance_fen += balance_fen
|
||
company_balances.append({
|
||
"company_id": m.company_id,
|
||
"company_name": company.name,
|
||
"balance": balance_fen,
|
||
"used": m.compute_balance_used or 0,
|
||
"role": m.role or "staff",
|
||
"joined_at": m.created_at,
|
||
})
|
||
|
||
# 从 compute-engine 获取用户实际额度(micro)
|
||
# quota = 个人充值 + 企业分配(已同步到引擎)
|
||
# 可用余额 = quota - used_quota
|
||
try:
|
||
from ...services import compute_client
|
||
engine_data = await compute_client.user_balance(username)
|
||
quota_micro = int(engine_data.get("quota", 0) or 0)
|
||
used_quota_micro = int(engine_data.get("used_quota", 0) or 0)
|
||
except Exception:
|
||
quota_micro = 0
|
||
used_quota_micro = 0
|
||
|
||
available_micro = max(0, quota_micro - used_quota_micro)
|
||
# 企业归集余额:分 → micro(1分=10000 micro)
|
||
total_company_balance_micro = total_company_balance_fen * 10000
|
||
# 个人余额 = users.compute_personal_balance(独立账本:个人充值/企业转入,单位分 → micro)
|
||
user_row = await db.session.get(User, user_id)
|
||
personal_balance_fen = (getattr(user_row, "compute_personal_balance", 0) or 0) if user_row else 0
|
||
personal_balance_micro = personal_balance_fen * 10000
|
||
|
||
return {
|
||
"personal_balance": personal_balance_micro,
|
||
"total_company_balance": total_company_balance_micro,
|
||
"total_balance": personal_balance_micro + total_company_balance_micro,
|
||
"quota": quota_micro,
|
||
"used_quota": used_quota_micro,
|
||
"available": available_micro,
|
||
"company_balances": company_balances,
|
||
}
|
||
|
||
|
||
@router.get("/user/usage", summary="个人端:使用记录")
|
||
async def user_usage(
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
user_id = user.get("id", "")
|
||
result = await db.session.execute(
|
||
select(ComputeUsageRecord).where(ComputeUsageRecord.user_id == user_id).order_by(ComputeUsageRecord.created_at.desc()).limit(100)
|
||
)
|
||
items = [{"id": r.id, "model": r.model, "token_count": r.token_count,
|
||
"standard_price": r.standard_price, "discount": r.discount,
|
||
"discount_source": r.discount_source, "actual_amount": r.actual_amount,
|
||
"balance_source": r.balance_source, "company_id": r.company_id,
|
||
"created_at": r.created_at}
|
||
for r in result.scalars().all()]
|
||
return {"items": items}
|
||
|
||
|
||
class LeaveCompanyBody(BaseModel):
|
||
reason: str = ""
|
||
|
||
|
||
@router.post("/user/companies/{company_id}/leave", summary="个人端:退出企业")
|
||
async def user_leave_company(
|
||
company_id: str,
|
||
body: LeaveCompanyBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
user_id = user.get("id", "")
|
||
|
||
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:
|
||
raise HTTPException(status_code=404, detail="您不是该企业成员")
|
||
|
||
if member.is_admin:
|
||
raise HTTPException(status_code=400, detail="企业管理员不可主动退出,请先转让管理员权限")
|
||
|
||
# 回收剩余企业分配的余额(仅 member.compute_balance 企业分配部分,绝不触碰个人充值余额)
|
||
remaining_balance = member.compute_balance or 0
|
||
if remaining_balance > 0:
|
||
company = await db.session.get(ParkCompany, company_id)
|
||
if company:
|
||
# 同步引擎 quota 扣减(分 → micro,1分=10000 micro)
|
||
from ...pay.service import _resolve_engine_user_id
|
||
from ...services import compute_client
|
||
total_micro = remaining_balance * MICRO_PER_FEN
|
||
try:
|
||
engine_user_id = await _resolve_engine_user_id(user.get("username", ""))
|
||
if engine_user_id:
|
||
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:
|
||
logger.warning("[leave_company] 引擎额度扣减失败 user=%s: %s", user_id, exc)
|
||
# 企业侧返还
|
||
company.compute_balance_used = max(0, (company.compute_balance_used or 0) - remaining_balance)
|
||
|
||
# 设置退出状态
|
||
member.status = "left"
|
||
member.left_at = now_str()
|
||
member.left_reason = body.reason
|
||
member.compute_balance = 0
|
||
|
||
await db.session.commit()
|
||
await write_audit(db, action="compute.user_leave_company", resource="company_member",
|
||
resource_id=member.id, detail=f"company={company_id}, refund={remaining_balance}",
|
||
user=user, request=request)
|
||
return {"ok": True, "refunded_balance": remaining_balance}
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
# 内部接口:价格查询、算力扣费
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
|
||
class PriceQueryBody(BaseModel):
|
||
model: str
|
||
token_count: int
|
||
|
||
|
||
@router.post("/price", summary="内部:查询实际价格")
|
||
async def compute_price(
|
||
body: PriceQueryBody,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
return await calculate_compute_price(db, user.get("id", ""), body.model, body.token_count)
|
||
|
||
|
||
class DeductBody(BaseModel):
|
||
model: str
|
||
token_count: int
|
||
|
||
|
||
@router.post("/deduct", summary="内部:算力扣费")
|
||
async def compute_deduct(
|
||
body: DeductBody,
|
||
request: Request,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
result = await deduct_compute_balance(db, user.get("id", ""), body.model, body.token_count)
|
||
if not result.get("success"):
|
||
raise HTTPException(status_code=402, detail=result.get("reason", "扣费失败"))
|
||
await write_audit(db, action="compute.deduct", resource="compute_usage_record",
|
||
resource_id="", detail=f"model={body.model}, tokens={body.token_count}, amount={result['amount']}",
|
||
user=user, request=request)
|
||
return result
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
# 企业充值订单(真实微信支付,复用 pay.service 下单/到账逻辑)
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
|
||
class CompanyRechargeOrderBody(BaseModel):
|
||
package_id: Optional[str] = None
|
||
amount_yuan: Optional[float] = None
|
||
client_type: str = "native"
|
||
|
||
|
||
@router.post("/enterprise/{company_id}/recharge/orders", summary="企业端:创建充值订单(微信支付)")
|
||
async def company_recharge_create(
|
||
company_id: str,
|
||
body: CompanyRechargeOrderBody,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
from ...pay import service as pay_service
|
||
await _get_admin_company(db, company_id, user)
|
||
try:
|
||
order = await pay_service.create_order(
|
||
db, user, client_type=body.client_type,
|
||
package_id=body.package_id or "", amount_yuan=body.amount_yuan,
|
||
target_type="company", target_id=company_id,
|
||
)
|
||
except ValueError as exc:
|
||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||
except RuntimeError as exc:
|
||
raise HTTPException(status_code=503, detail=str(exc)) from exc
|
||
return order
|
||
|
||
|
||
@router.post("/enterprise/{company_id}/recharge/orders/{order_no}/status", summary="企业端:查询充值订单状态")
|
||
async def company_recharge_status(
|
||
company_id: str,
|
||
order_no: str,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
from ...pay import service as pay_service
|
||
await _get_admin_company(db, company_id, user)
|
||
repo = db.compute_recharges
|
||
order = await repo.get_by_order_no(order_no)
|
||
if order is None:
|
||
raise HTTPException(status_code=404, detail="订单不存在")
|
||
if order.get("target_type") != "company" or order.get("target_id") != company_id:
|
||
raise HTTPException(status_code=403, detail="订单不属于该企业")
|
||
order = await pay_service.reconcile(db, order)
|
||
return {
|
||
"order_no": order["order_no"], "status": order["status"],
|
||
"paid": order["status"] in ("paid", "credited"),
|
||
"credited": order["status"] == "credited",
|
||
"amount_fen": order["amount_fen"], "quota_micro": order["quota_micro"],
|
||
"transaction_id": order.get("transaction_id", ""), "expires_at": order["expires_at"],
|
||
"paid_at": order.get("paid_at", ""), "credited_at": order.get("credited_at", ""),
|
||
}
|
||
|
||
|
||
@router.get("/enterprise/{company_id}/recharge/orders", summary="企业端:充值订单记录")
|
||
async def company_recharge_orders(
|
||
company_id: str,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
):
|
||
await _get_admin_company(db, company_id, user)
|
||
rows = (await db.session.scalars(
|
||
select(ComputeRecharge)
|
||
.where(ComputeRecharge.company_id == company_id)
|
||
.order_by(ComputeRecharge.created_at.desc())
|
||
.limit(50)
|
||
)).all()
|
||
return {"items": [r.__dict__ for r in rows]}
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
# 企业↔个人余额互转
|
||
# ═══════════════════════════════════════════════════════════════════
|
||
|
||
class TransferBody(BaseModel):
|
||
amount_yuan: float
|
||
reason: Optional[str] = ""
|
||
|
||
|
||
MICRO_PER_FEN = 10000 # 1分 = 10000 micro(1元=100分=1,000,000 micro)
|
||
|
||
|
||
@router.post("/enterprise/{company_id}/transfer/to-personal", summary="企业端:企业余额转个人")
|
||
async def transfer_to_personal(
|
||
company_id: str,
|
||
body: TransferBody,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
request: Request = None,
|
||
):
|
||
from ...pay.service import _resolve_engine_user_id
|
||
from ...services import compute_client
|
||
|
||
company = await _get_admin_company(db, company_id, user)
|
||
amount_fen = int(round(body.amount_yuan * 100))
|
||
if amount_fen <= 0:
|
||
raise HTTPException(status_code=400, detail="金额必须大于0")
|
||
|
||
# ★ 企业可用余额 = 总额 - 已使用(已分配给成员/转出的部分)
|
||
company_available = max(0, (company.compute_balance or 0) - (company.compute_balance_used or 0))
|
||
if company_available < amount_fen:
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail=f"企业可用余额不足:可用 ¥{company_available / 100:.2f},尝试转出 ¥{body.amount_yuan:.2f}",
|
||
)
|
||
|
||
# 增加个人 quota
|
||
username = user.get("username", "")
|
||
engine_user_id = await _resolve_engine_user_id(username)
|
||
if not engine_user_id:
|
||
raise HTTPException(status_code=500, detail="个人算力账号未就绪")
|
||
total_micro = amount_fen * MICRO_PER_FEN
|
||
try:
|
||
await compute_client.adjust_user_quota(engine_user_id, total_micro, "add")
|
||
except Exception as exc:
|
||
raise HTTPException(status_code=400, detail=f"个人余额增加失败: {exc}") from exc
|
||
await compute_client.sync_user_mirror(db, engine_user_id)
|
||
|
||
# 扣减企业可用余额(增加已使用额,不减少总额)
|
||
company.compute_balance_used = (company.compute_balance_used or 0) + amount_fen
|
||
|
||
# 个人余额账本同步(分):企业转入 → users.compute_personal_balance
|
||
from ..infrastructure.models import User as _User
|
||
from sqlalchemy import update as _sa_update
|
||
await db.session.execute(
|
||
_sa_update(_User).where(_User.id == user["id"]).values(
|
||
compute_personal_balance=_User.compute_personal_balance + amount_fen
|
||
)
|
||
)
|
||
|
||
# 写入个人账单(企业转入)
|
||
await _write_transfer_bill(
|
||
db, target_type="user", target_id=user["id"], amount_fen=amount_fen,
|
||
source_label="企业转入", operator_id=user["id"], from_name=company.name,
|
||
)
|
||
|
||
await db.session.commit()
|
||
|
||
await write_audit(db, action="compute.transfer_to_personal", resource="park_company",
|
||
resource_id=company_id, detail=f"{amount_fen} fen → user {user['id']}, {total_micro} micro, reason={body.reason}",
|
||
user=user, request=request)
|
||
return {"ok": True, "amount_fen": amount_fen, "company_balance": company.compute_balance,
|
||
"company_balance_used": company.compute_balance_used,
|
||
"company_available": max(0, (company.compute_balance or 0) - (company.compute_balance_used or 0))}
|
||
|
||
|
||
@router.post("/enterprise/{company_id}/transfer/from-personal", summary="企业端:个人余额转企业")
|
||
async def transfer_from_personal(
|
||
company_id: str,
|
||
body: TransferBody,
|
||
db: Database = Depends(get_db),
|
||
user: dict = Depends(get_current_user),
|
||
request: Request = None,
|
||
):
|
||
from ...pay.service import _resolve_engine_user_id
|
||
from ...services import compute_client
|
||
|
||
company = await _get_admin_company(db, company_id, user)
|
||
amount_fen = int(round(body.amount_yuan * 100))
|
||
if amount_fen <= 0:
|
||
raise HTTPException(status_code=400, detail="金额必须大于0")
|
||
|
||
# 扣减个人 quota
|
||
username = user.get("username", "")
|
||
engine_user_id = await _resolve_engine_user_id(username)
|
||
if not engine_user_id:
|
||
raise HTTPException(status_code=500, detail="个人算力账号未就绪")
|
||
total_micro = amount_fen * MICRO_PER_FEN
|
||
|
||
# ★ 服务端严格校验个人可用余额(以平台账本为权威,引擎余额为参考)
|
||
from ..infrastructure.models import User as _User
|
||
from sqlalchemy import update as _sa_update
|
||
user_row = await db.session.get(_User, user["id"])
|
||
ledger_balance = (getattr(user_row, "compute_personal_balance", 0) or 0) if user_row else 0
|
||
try:
|
||
user_bal = await compute_client.user_balance(username)
|
||
quota_micro = int(user_bal.get("quota", 0) or 0)
|
||
used_micro = int(user_bal.get("used_quota", 0) or 0)
|
||
available_micro = max(0, quota_micro - used_micro)
|
||
except Exception as exc:
|
||
raise HTTPException(status_code=502, detail=f"查询个人余额失败: {exc}") from exc
|
||
if ledger_balance < amount_fen:
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail=f"个人可转余额不足:账本可用 ¥{ledger_balance / 100:.2f},尝试转出 ¥{body.amount_yuan:.2f}",
|
||
)
|
||
if available_micro < total_micro:
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail=f"个人可用余额不足:可用 ¥{available_micro / 1000000:.2f},尝试转出 ¥{body.amount_yuan:.2f}",
|
||
)
|
||
|
||
# 先扣账本(幂等条件更新,防止并发超扣)
|
||
upd = await db.session.execute(
|
||
_sa_update(_User)
|
||
.where(_User.id == user["id"], _User.compute_personal_balance >= amount_fen)
|
||
.values(compute_personal_balance=_User.compute_personal_balance - amount_fen)
|
||
)
|
||
if upd.rowcount == 0:
|
||
raise HTTPException(status_code=400, detail="个人可转余额不足(账本校验失败)")
|
||
|
||
try:
|
||
await compute_client.adjust_user_quota(engine_user_id, total_micro, "subtract")
|
||
except Exception as exc:
|
||
raise HTTPException(status_code=400, detail=f"个人余额扣减失败: {exc}") from exc
|
||
await compute_client.sync_user_mirror(db, engine_user_id)
|
||
|
||
# 增加企业余额
|
||
company.compute_balance = (company.compute_balance or 0) + amount_fen
|
||
|
||
# 写入企业账单(个人转入)
|
||
await _write_transfer_bill(
|
||
db, target_type="company", target_id=company_id, amount_fen=amount_fen,
|
||
source_label="个人转入", operator_id=user["id"], from_name=user.get("nickname") or user.get("username", ""),
|
||
)
|
||
|
||
await db.session.commit()
|
||
|
||
await write_audit(db, action="compute.transfer_from_personal", resource="park_company",
|
||
resource_id=company_id, detail=f"user {user['id']} {total_micro} micro → {amount_fen} fen, reason={body.reason}",
|
||
user=user, request=request)
|
||
return {"ok": True, "amount_fen": amount_fen, "company_balance": company.compute_balance}
|