Files
taiji-AI-PAD/services/mcp-server/app/routes/billing_webhook.py
T
2026-01-16 10:38:17 +00:00

556 lines
22 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 的运行结束通知
- EU计算:基于Pod运行时长(调用 app.billing.calculate_eu)
- 公式:EU = ceil(duration_seconds / 10)
- 存储表:agent_billing_records
- 注意:周期性计费任务会预先扣款,此回调只扣除增量部分
"""
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
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)}")
logger.info(f"📥 收到LiteLLM回调 - 类型: {'数组' if isinstance(body, list) else '对象'}")
# 判断是数组还是单个对象
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 的计费记录:
- 如果存在运行中的计费记录(end_time=None),则更新该记录
- 如果不存在运行中的记录,则创建新记录(异常情况的兜底)
计费逻辑:
- 周期计费任务会对运行中的 Agent 进行增量扣款
- 此回调负责结算最终费用,只扣除增量部分,避免重复扣款
"""
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}")
# 计算成本和EU
duration_seconds = callback_data.podRunningTimeSeconds
eu_consumed = calculate_eu(duration_seconds)
# 假设是平台Agent(可以根据agent_name前缀判断)
is_platform_agent = True
agent_type = "platform" # 可以从agent名称中提取
new_cost = calculate_platform_agent_cost(agent_type, duration_seconds)
# ✅ 先查找现有的运行中计费记录(end_time == None)
existing_result = await db.execute(
select(AgentBillingRecord).where(
and_(
AgentBillingRecord.agent_name == callback_data.agentName,
AgentBillingRecord.user_id == callback_data.userId,
AgentBillingRecord.end_time == None # 运行中的记录
)
)
)
existing_record = existing_result.scalar_one_or_none()
if existing_record:
# ✅ 更新现有记录(避免重复创建)
previous_cost = Decimal(str(existing_record.cost or 0))
existing_record.end_time = end_time or datetime.utcnow()
existing_record.duration_seconds = duration_seconds
existing_record.eu_consumed = eu_consumed
existing_record.cost = float(new_cost)
existing_record.period_end = end_time or datetime.utcnow()
existing_record.tools_used = callback_data.toolsUsed
existing_record.request_id = callback_data.requestId
billing_record = existing_record
# 计算增量成本(新成本 - 已扣成本)
cost_increment = new_cost - previous_cost
logger.info(
f"📝 更新现有计费记录: agent={callback_data.agentName}, "
f"之前成本={previous_cost}, 最终成本={new_cost}, 增量={cost_increment}"
)
# 只扣除增量部分(避免与周期计费重复扣款)
if cost_increment > 0:
success, message = await deduct_balance(
callback_data.userId,
cost_increment,
db,
f"Agent 结算(增量): {callback_data.agentName}"
)
if not success:
logger.warning(f"增量余额扣除失败: {message}")
else:
logger.info(f"无需扣款(增量={cost_increment})")
else:
# ⚠️ 没有现有记录,创建新记录(异常情况的兜底)
logger.warning(
f"⚠️ 未找到运行中的计费记录,将创建新记录: "
f"agent={callback_data.agentName}, user={callback_data.userId}"
)
billing_record = AgentBillingRecord(
user_id=callback_data.userId,
channel_id=user.channel_id if user.channel_id else None,
agent_name=callback_data.agentName,
agent_type=agent_type,
is_platform_agent=is_platform_agent,
duration_seconds=duration_seconds,
eu_consumed=eu_consumed,
cost=float(new_cost),
start_time=start_time or datetime.utcnow(),
end_time=end_time or datetime.utcnow(),
period_start=start_time or datetime.utcnow(),
period_end=end_time or datetime.utcnow(),
tools_used=callback_data.toolsUsed,
request_id=callback_data.requestId,
)
db.add(billing_record)
# 新记录需要全额扣款
success, message = await deduct_balance(
callback_data.userId,
new_cost,
db,
f"Agent 运行: {callback_data.agentName}"
)
if not success:
logger.warning(f"余额扣除失败: {message}")
await db.commit()
await db.refresh(billing_record)
logger.info(
f"✅ Agent 计费记录{'更新' if existing_record else '创建'}成功: "
f"agent={callback_data.agentName}, duration={duration_seconds}秒, "
f"EU={eu_consumed}, cost={new_cost}, tools={callback_data.toolsUsed}"
)
return AgentManagerCallbackResponse(
success=True,
message=f"Agent 计费记录{'更新' if existing_record else '创建'}成功",
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"}