Files
taiji-AI-PAD/services/mcp-server/app/routes/billing_webhook.py
T
2026-03-13 02:53:27 +00:00

558 lines
23 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
LiteLLM Callback Webhook路由
接收LiteLLM的实时Token计费数据
此模块处理两种计费回调:
1. LiteLLM 模型调用回调(/litellm-callback)
- 数据来源:LiteLLM 的 success_callback
- EU计算:基于Token数量和模型类型
- 公式:EU = total_tokens * MODEL_EU_RATE[model_name]
- 存储表:model_billing_records
- 扣款时机:立即扣款
2. Agent Manager 回调(/agent-callback)
- 数据来源:Agent Manager 的 API 调用通知
- 计费类型:API 调用计费(record_type="api_call")
- 计费标准:每次调用固定 0.001 EU(= $0.001),与运行时间无关
- 公式:Cost = 0.001 USD/call,EU = Cost
- 存储表:agent_billing_records
- 扣款时机:每次 API 调用立即扣款
- 注意:这与 VM 运行时间计费(record_type="vm_runtime")不同,
VM 运行时间计费按小时费率由周期任务处理
"""
import json
import logging
from datetime import datetime
from typing import Optional, Dict, Any, List, Union
from decimal import Decimal
from fastapi import APIRouter, Depends, HTTPException, Request
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, update
from pydantic import BaseModel, Field
from database import get_db
from models import ModelBillingRecord, TenantModelKey, Balance, User
from app.schemas import AgentManagerCallbackData, AgentManagerCallbackResponse
from app.db_utils import ensure_idempotent, atomic_eu_consume, with_retry
from app.billing import calculate_api_call_cost
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/v1/billing", tags=["计费Webhook"])
class LiteLLMUsage(BaseModel):
"""Token使用量"""
prompt_tokens: int = 0
completion_tokens: int = 0
total_tokens: int = 0
class LiteLLMCallbackData(BaseModel):
"""LiteLLM Callback数据 - 支持LiteLLM实际发送的格式"""
call_id: str = Field(default="", alias="id")
trace_id: Optional[str] = None
model: str = ""
call_type: Optional[str] = None
cache_hit: Optional[bool] = None
stream: Optional[bool] = None
status: str = "success"
custom_llm_provider: Optional[str] = None
startTime: Optional[Union[str, float]] = None # 可能是字符串或Unix时间戳
endTime: Optional[Union[str, float]] = None
response_time: Optional[float] = None
response_cost: Optional[float] = None
total_tokens: Optional[int] = None
prompt_tokens: Optional[int] = None
completion_tokens: Optional[int] = None
api_key: Optional[str] = None
team_id: Optional[str] = None
api_base: Optional[str] = None
model_group: Optional[str] = None
model_id: Optional[str] = None
messages: Optional[List[Dict[str, Any]]] = None
response: Optional[Dict[str, Any]] = None
metadata: Optional[Dict[str, Any]] = None
hidden_params: Optional[Dict[str, Any]] = None
model_map_information: Optional[Dict[str, Any]] = None
cost_breakdown: Optional[Dict[str, Any]] = None
error_str: Optional[str] = None
error_information: Optional[Dict[str, Any]] = None
class Config:
populate_by_name = True
extra = "allow" # 允许额外字段
async def get_tenant_from_api_key(api_key_hash: str, db: AsyncSession) -> tuple[Optional[str], Optional[str]]:
"""从API Key Hash解析租户ID和渠道ID"""
if not api_key_hash:
logger.debug("API Key Hash为空")
return None, None
try:
# 打印接收到的api_key_hash用于调试
logger.info(f"🔍 查询租户信息 - api_key_hash: {api_key_hash[:20] if len(api_key_hash) > 20 else api_key_hash}...")
# 方式1: 直接匹配 litellm_key_hash (加密存储的key)
result = await db.execute(
select(TenantModelKey).where(TenantModelKey.litellm_key_hash == api_key_hash)
)
key_record = result.scalar_one_or_none()
if key_record:
logger.info(f"✅ 通过litellm_key_hash找到租户: {key_record.tenant_id}")
return str(key_record.tenant_id), str(key_record.channel_id) if key_record.channel_id else None
# 方式2: 尝试匹配 litellm_key_id (LiteLLM返回的完整key)
result = await db.execute(
select(TenantModelKey).where(TenantModelKey.litellm_key_id == api_key_hash)
)
key_record = result.scalar_one_or_none()
if key_record:
logger.info(f"✅ 通过litellm_key_id找到租户: {key_record.tenant_id}")
return str(key_record.tenant_id), str(key_record.channel_id) if key_record.channel_id else None
# 方式3: 模糊匹配 (如果api_key_hash是截断的)
if len(api_key_hash) > 10:
result = await db.execute(
select(TenantModelKey).where(
TenantModelKey.litellm_key_id.like(f"{api_key_hash[:10]}%")
)
)
key_record = result.scalar_one_or_none()
if key_record:
logger.info(f"✅ 通过前缀匹配找到租户: {key_record.tenant_id}")
return str(key_record.tenant_id), str(key_record.channel_id) if key_record.channel_id else None
# 调试:列出所有key的前缀
result = await db.execute(select(TenantModelKey).limit(10))
all_keys = result.scalars().all()
logger.warning(f"❌ 未找到匹配的租户。数据库中的key前缀: {[k.litellm_key_id[:20] if k.litellm_key_id else 'None' for k in all_keys]}")
except Exception as e:
logger.error(f"❌ 查询租户信息失败: {e}", exc_info=True)
return None, None
def calculate_eu_from_tokens(total_tokens: int, model_name: str) -> Decimal:
"""根据Token计算EU消耗"""
# 不同模型的Token到EU转换率
MODEL_EU_RATE = {
"gpt-4": Decimal("0.0001"), # 1000 tokens = 0.1 EU
"gpt-3.5-turbo": Decimal("0.00005"), # 1000 tokens = 0.05 EU
"claude": Decimal("0.0001"),
"default": Decimal("0.0001"),
}
# 匹配模型名称
rate = MODEL_EU_RATE["default"]
for model_key, model_rate in MODEL_EU_RATE.items():
if model_key in model_name.lower():
rate = model_rate
break
return Decimal(total_tokens) * rate
async def process_single_callback(callback_data: LiteLLMCallbackData, db: AsyncSession) -> Dict[str, Any]:
"""处理单个LiteLLM回调数据"""
call_id = callback_data.call_id or callback_data.trace_id or ""
if not call_id:
logger.warning("回调数据缺少call_id,跳过")
return {"message": "Skipped - no call_id", "call_id": None}
# 检查是否已处理过(幂等性)
existing = await db.execute(
select(ModelBillingRecord).where(
ModelBillingRecord.litellm_call_id == call_id
)
)
if existing.scalar_one_or_none():
logger.info(f"LiteLLM回调已处理过: {call_id}")
return {"message": "Already processed", "call_id": call_id}
# 从metadata中提取信息
api_key = None
team_id = callback_data.team_id
tenant_id = None
channel_id = None
if callback_data.metadata:
api_key = callback_data.metadata.get("user_api_key_hash")
if not team_id:
team_id = callback_data.metadata.get("user_api_key_team_id")
# ✅ 优先从 user_api_key_auth_metadata 直接获取租户信息(LiteLLM创建key时设置的metadata)
auth_metadata = callback_data.metadata.get("user_api_key_auth_metadata")
if auth_metadata and isinstance(auth_metadata, dict):
tenant_id = auth_metadata.get("tenant_id")
channel_id = auth_metadata.get("channel_id")
logger.info(f"✅ 从auth_metadata获取租户信息: tenant_id={tenant_id}, channel_id={channel_id}, tenant_name={auth_metadata.get('tenant_name')}")
# 调试日志
logger.info(f"📦 Callback metadata: api_key_hash={api_key[:20] if api_key else 'None'}..., team_id={team_id}, model={callback_data.model}")
else:
logger.warning("⚠️ Callback数据缺少metadata")
# 如果auth_metadata中没有租户信息,尝试从数据库查询(兜底方案)
if not tenant_id and api_key:
tenant_id, channel_id = await get_tenant_from_api_key(api_key, db)
if not tenant_id:
logger.warning(f"⚠️ 无法解析租户ID: api_key={api_key[:10] if api_key else 'None'}..., call_id={call_id}")
else:
logger.info(f"✅ 成功解析租户: tenant_id={tenant_id}, channel_id={channel_id}")
# 计算Token - 使用顶级字段或从metadata中获取
total_tokens = callback_data.total_tokens or 0
prompt_tokens = callback_data.prompt_tokens or 0
completion_tokens = callback_data.completion_tokens or 0
if callback_data.metadata and callback_data.metadata.get("usage_object"):
usage_obj = callback_data.metadata["usage_object"]
total_tokens = total_tokens or usage_obj.get("total_tokens", 0)
prompt_tokens = prompt_tokens or usage_obj.get("prompt_tokens", 0)
completion_tokens = completion_tokens or usage_obj.get("completion_tokens", 0)
eu_consumed = calculate_eu_from_tokens(total_tokens, callback_data.model)
# 解析时间 - 支持Unix时间戳和ISO字符串
start_time = None
end_time = None
try:
if callback_data.startTime:
if isinstance(callback_data.startTime, (int, float)):
start_time = datetime.fromtimestamp(callback_data.startTime)
else:
dt = datetime.fromisoformat(str(callback_data.startTime).replace('Z', '+00:00'))
start_time = dt.replace(tzinfo=None)
if callback_data.endTime:
if isinstance(callback_data.endTime, (int, float)):
end_time = datetime.fromtimestamp(callback_data.endTime)
else:
dt = datetime.fromisoformat(str(callback_data.endTime).replace('Z', '+00:00'))
end_time = dt.replace(tzinfo=None)
except Exception as e:
logger.warning(f"时间解析失败: {e}")
# ✅ 原子操作:同时创建记录和更新余额
if not tenant_id:
logger.error(f"无法解析租户ID,跳过计费: call_id={call_id}")
raise HTTPException(status_code=400, detail="无法解析租户ID")
try:
# 创建计费记录
record = ModelBillingRecord(
tenant_id=tenant_id,
channel_id=channel_id,
litellm_call_id=call_id,
api_key=api_key[:8] + "..." + api_key[-4:] if api_key and len(api_key) > 12 else api_key,
team_id=team_id,
model_name=callback_data.model,
input_tokens=prompt_tokens,
output_tokens=completion_tokens,
total_tokens=total_tokens,
total_cost=Decimal(callback_data.response_cost or 0),
eu_consumed=eu_consumed,
status=callback_data.status,
start_time=start_time,
end_time=end_time,
response_time_ms=int((callback_data.response_time or 0) * 1000),
raw_callback_data=callback_data.model_dump() if hasattr(callback_data, 'model_dump') else callback_data.dict()
)
db.add(record)
# 使用行锁保护余额更新,防止并发扣款
balance_result = await db.execute(
select(Balance)
.where(Balance.user_id == tenant_id)
.with_for_update() # 行锁
)
balance = balance_result.scalar_one_or_none()
if balance is None:
# 如果余额记录不存在,创建一个新的(初始余额为0)
balance = Balance(user_id=tenant_id, eu_balance=0)
db.add(balance)
await db.flush() # 确保记录创建
logger.warning(f"用户 {tenant_id} 余额记录不存在,已创建初始余额为0")
# 扣减EU余额(使用Decimal精确计算)
old_balance = Decimal(str(balance.eu_balance))
new_balance = old_balance - Decimal(str(eu_consumed))
balance.eu_balance = float(new_balance)
# 使用原子操作更新用户表的 total_eu_consumed 字段
await db.execute(
update(User)
.where(User.id == tenant_id)
.values(total_eu_consumed=User.total_eu_consumed + float(eu_consumed))
)
logger.info(
f"✅ 扣减用户余额: user_id={tenant_id}, "
f"eu_consumed={eu_consumed}, "
f"old_balance={old_balance:.4f}, "
f"new_balance={new_balance:.4f}"
)
# 余额不足警告(但不阻止记录)
if new_balance < 0:
logger.warning(
f"⚠️ 用户余额不足: user_id={tenant_id}, "
f"balance={new_balance:.4f}, "
f"建议充值"
)
except Exception as e:
logger.error(f"❌ 计费处理失败: {e}", exc_info=True)
# 任何异常都应该回滚整个事务
raise
logger.info(
f"记录模型调用计费: {callback_data.model}, "
f"tokens={total_tokens}, EU={eu_consumed}, "
f"tenant_id={tenant_id}"
)
return {
"message": "Success",
"call_id": call_id,
"record_id": str(record.id) if record.id else None,
"eu_consumed": float(eu_consumed),
"balance_updated": tenant_id is not None
}
@router.post("/litellm-callback")
async def litellm_callback(
request: Request,
db: AsyncSession = Depends(get_db)
):
"""接收LiteLLM的Token计费回调 - 支持单个对象或数组"""
try:
body = await request.json()
# 🔍 调试:打印完整的回调内容(使用print确保输出)
print(f"📥 收到LiteLLM回调 - 类型: {'数组' if isinstance(body, list) else '对象'}")
print(f"📥 完整回调内容:\n{json.dumps(body, indent=2, ensure_ascii=False)}")
# 判断是数组还是单个对象
if isinstance(body, list):
# 处理数组(批量回调)
results = []
for item in body:
try:
callback_data = LiteLLMCallbackData(**item)
result = await process_single_callback(callback_data, db)
results.append(result)
except Exception as e:
logger.error(f"处理单个回调失败: {e}")
results.append({"message": f"Error: {str(e)}", "call_id": item.get("id")})
await db.commit()
return {"message": "Batch processed", "count": len(results), "results": results}
else:
# 处理单个对象
callback_data = LiteLLMCallbackData(**body)
result = await process_single_callback(callback_data, db)
await db.commit()
return result
except Exception as e:
logger.error(f"LiteLLM Callback处理失败: {e}", exc_info=True)
await db.rollback()
raise HTTPException(status_code=500, detail=str(e))
@router.get("/litellm-callback/health")
async def callback_health():
"""Webhook健康检查"""
return {"status": "ok", "endpoint": "/api/v1/billing/litellm-callback"}
# ============= Agent Manager 回调接口 =============
@router.post("/agent-callback")
async def agent_manager_callback(
request: Request,
db: AsyncSession = Depends(get_db)
):
"""
接收 Agent Manager 的回调数据
Agent Manager 在用户调用 agent 后返回运行时信息:
- Pod 运行时间(记录但不用于计费)
- 使用的工具列表(可选)
此接口用于创建 Agent 的 API 调用计费记录:
- record_type="api_call"(API 调用计费)
- 每次调用固定费用:0.001 EU(= $0.001)
- 与运行时间无关,按调用次数计费
计费逻辑:
- API 调用计费:每次调用立即扣款 0.001 EU
- VM 运行时间计费(record_type="vm_runtime"):由周期计费任务处理,按小时费率计费
注意:API 调用计费与 VM 运行时间计费是两种不同的计费方式
"""
from models import AgentBillingRecord
from sqlalchemy import and_
from app.billing import calculate_eu, calculate_platform_agent_cost, deduct_balance
try:
# 解析请求体
body = await request.json()
logger.info(f"📥 收到 Agent Manager 回调: {json.dumps(body, indent=2, ensure_ascii=False)}")
callback_data = AgentManagerCallbackData(**body)
# 验证用户存在
result = await db.execute(
select(User).where(User.id == callback_data.userId)
)
user = result.scalar_one_or_none()
if not user:
logger.warning(f"用户不存在: {callback_data.userId}")
raise HTTPException(status_code=404, detail=f"用户不存在: {callback_data.userId}")
# 解析时间
start_time = None
end_time = None
if callback_data.startTime:
try:
start_time = datetime.fromisoformat(callback_data.startTime.replace('Z', '+00:00'))
if start_time.tzinfo:
start_time = start_time.replace(tzinfo=None)
except Exception as e:
logger.warning(f"开始时间解析失败: {e}")
if callback_data.endTime:
try:
end_time = datetime.fromisoformat(callback_data.endTime.replace('Z', '+00:00'))
if end_time.tzinfo:
end_time = end_time.replace(tzinfo=None)
except Exception as e:
logger.warning(f"结束时间解析失败: {e}")
# ✅ 验证必填字段,防止空值
agent_name = callback_data.agentName or "unknown"
if not callback_data.agentName:
logger.warning(
f"⚠️ Agent Manager回调缺少agentName: "
f"user_id={callback_data.userId}, request_id={callback_data.requestId}"
)
# ✅ 计算 API 调用成本:每次调用固定 0.001 EU(= $0.001)
# API 调用计费与运行时间无关,每次调用固定费用
duration_seconds = callback_data.podRunningTimeSeconds # 仅记录,不用于计费
# 假设是平台Agent(可以从agent_name前缀判断)
is_platform_agent = True
agent_type = "platform" # 可以从agent名称中提取
# ✅ 使用 API 调用计费函数:固定 0.001 EU/call
cost = calculate_api_call_cost()
# 调试日志 - 使用 print 确保输出
print(f"\n⭐️⭐️⭐️ Agent Manager 回调 ⭐️⭐️⭐️")
print(f"User: {callback_data.userId}")
print(f"Agent: {agent_name}")
print(f"Cost: {cost} (type: {type(cost)})")
print(f"Duration: {duration_seconds}秒")
print("⭐️⭐️⭐️⭐️⭐️⭐️⭐️⭐️⭐️⭐️⭐️\n", flush=True)
logger.info(f"📝 Agent Manager 回调收到: user={callback_data.userId}, agent={agent_name}, cost={cost}, type={type(cost)}")
# EU = Cost(1 EU = 1 美元)
eu_consumed = float(cost)
# ✅ 新逻辑:每次调用都创建新的 API 调用记录
# 这是 API 调用计费,不是 VM 运行时间计费
# VM 运行时间计费由周期任务处理
billing_record = AgentBillingRecord(
record_type="api_call", # ✅ API 调用计费
user_id=callback_data.userId,
channel_id=user.channel_id if user.channel_id else None,
agent_name=agent_name,
agent_type=agent_type,
is_platform_agent=is_platform_agent,
duration_seconds=duration_seconds, # 本次调用的处理时间
request_count=1, # 这是一次调用
eu_consumed=eu_consumed,
cost=float(cost),
start_time=start_time or datetime.utcnow(),
end_time=end_time or datetime.utcnow(), # API 调用已完成
period_start=start_time or datetime.utcnow(),
period_end=end_time or datetime.utcnow(),
tools_used=callback_data.toolsUsed, # 本次调用使用的工具
request_id=callback_data.requestId, # 本次调用的唯一ID
)
db.add(billing_record)
# ✅ API 调用立即扣款
print(f"\n🔍 准备扣款: cost={cost}, type={type(cost)}, cost>0={cost > 0}", flush=True)
logger.info(f"🔍 准备扣款: cost={cost}, type={type(cost)}, cost>0={cost > 0}")
if cost > 0:
print(f"✅ Cost > 0, 开始调用 deduct_balance", flush=True)
logger.info(f"🔍 开始调用 deduct_balance: user_id={callback_data.userId}, cost={cost}")
success, message = await deduct_balance(
callback_data.userId, cost, db,
f"Agent API调用: {agent_name}",
auto_commit=False # 🔧 修复:与billing_record在同一事务中
)
print(f"✅ deduct_balance 返回: success={success}, message={message}", flush=True)
logger.info(f"🔍 deduct_balance 返回: success={success}, message={message}")
if success:
logger.info(
f"💰 Agent API调用计费成功: agent={agent_name}, "
f"request_id={callback_data.requestId}, duration={duration_seconds}秒, "
f"tools={callback_data.toolsUsed}, cost={cost}"
)
else:
logger.warning(
f"⚠️ Agent API调用扣款失败: agent={agent_name}, "
f"request_id={callback_data.requestId}, cost={cost}, 原因={message}"
)
else:
logger.info(
f"📝 创建Agent API调用记录(费用为0): agent={agent_name}, "
f"request_id={callback_data.requestId}"
)
print(f"\n🔥 准备 commit...", flush=True)
await db.commit()
print(f"🔥 commit 完成!\n", flush=True)
await db.refresh(billing_record)
logger.info(
f"✅ Agent API调用计费记录创建成功: "
f"agent={agent_name}, request_id={callback_data.requestId}, "
f"duration={duration_seconds}秒, EU={eu_consumed}, cost={cost}, "
f"tools={callback_data.toolsUsed}"
)
return AgentManagerCallbackResponse(
success=True,
message="Agent API调用计费记录创建成功",
recordId=str(billing_record.id)
)
except HTTPException:
raise
except Exception as e:
logger.error(f"Agent Manager 回调处理失败: {e}", exc_info=True)
await db.rollback()
raise HTTPException(status_code=500, detail=f"回调处理失败: {str(e)}")
@router.get("/agent-callback/health")
async def agent_callback_health():
"""Agent回调健康检查"""
return {"status": "ok", "endpoint": "/api/v1/billing/agent-callback"}