feat: add ai cost tracking foundation

This commit is contained in:
2026-07-31 16:24:35 +08:00
parent ec9f8a015a
commit a84650bb9e
14 changed files with 327 additions and 2 deletions

View File

@@ -133,6 +133,12 @@ class AdminDashboardService:
msg_filter.append(ChatMessage.created_at <= end)
ai_filter.append(AiRequestLog.created_at <= end)
cost_currency = db.scalar(
select(AiRequestLog.currency)
.where(and_(*ai_filter), AiRequestLog.currency.is_not(None))
.order_by(AiRequestLog.id.desc())
.limit(1)
)
return {
"userCount": db.scalar(select(func.count(User.id)).where(and_(*user_filter))) or 0,
"sessionCount": db.scalar(select(func.count(ChatSession.id)).where(and_(*session_filter))) or 0,
@@ -142,6 +148,8 @@ class AdminDashboardService:
"inputToken": db.scalar(select(func.coalesce(func.sum(AiRequestLog.input_token), 0)).where(and_(*ai_filter))) or 0,
"outputToken": db.scalar(select(func.coalesce(func.sum(AiRequestLog.output_token), 0)).where(and_(*ai_filter))) or 0,
"totalToken": db.scalar(select(func.coalesce(func.sum(AiRequestLog.total_token), 0)).where(and_(*ai_filter))) or 0,
"estimatedCost": float(db.scalar(select(func.coalesce(func.sum(AiRequestLog.estimated_cost), 0)).where(and_(*ai_filter))) or 0),
"costCurrency": cost_currency,
}

View File

@@ -2,9 +2,11 @@ from __future__ import annotations
import json
from collections.abc import Sequence
from decimal import Decimal, ROUND_HALF_UP
from sqlalchemy.orm import Session
from app.models.ai_config import ModelConfig
from app.models.logs import AiRequestLog
from app.services.rag_service import RetrievedChunk
@@ -25,12 +27,18 @@ class AiRequestLogService:
output_token: int,
cost_ms: int,
retrieved_chunks: Sequence[RetrievedChunk] | None = None,
model_id: int | None = None,
route_reason: str | None = None,
question_type: str | None = None,
) -> None:
model = db.get(ModelConfig, model_id) if model_id else None
estimated_cost, currency = _estimate_cost(model, input_token, output_token)
db.add(
AiRequestLog(
session_id=session_id,
message_id=message_id,
user_id=user_id,
model_id=model_id,
model_name=model_name,
prompt=prompt,
knowledge_ids=knowledge_ids,
@@ -40,6 +48,11 @@ class AiRequestLogService:
output_token=output_token,
total_token=input_token + output_token,
cost_ms=cost_ms,
estimated_cost=estimated_cost,
currency=currency,
route_reason=route_reason or _default_route_reason(retrieve_count),
question_type=question_type or _question_type(retrieve_count),
knowledge_hit=1 if retrieve_count > 0 else 0,
status="SUCCESS",
)
)
@@ -58,18 +71,25 @@ class AiRequestLogService:
cost_ms: int,
error_message: str,
retrieved_chunks: Sequence[RetrievedChunk] | None = None,
model_id: int | None = None,
route_reason: str | None = None,
question_type: str | None = None,
) -> None:
db.add(
AiRequestLog(
session_id=session_id,
message_id=message_id,
user_id=user_id,
model_id=model_id,
model_name=model_name,
prompt=prompt,
knowledge_ids=knowledge_ids,
retrieve_count=retrieve_count,
retrieved_chunks=_dump_retrieved_chunks(retrieved_chunks),
cost_ms=cost_ms,
route_reason=route_reason or _default_route_reason(retrieve_count),
question_type=question_type or _question_type(retrieve_count),
knowledge_hit=1 if retrieve_count > 0 else 0,
status="FAILED",
error_message=error_message,
)
@@ -91,3 +111,22 @@ def _dump_retrieved_chunks(chunks: Sequence[RetrievedChunk] | None) -> str | Non
for index, chunk in enumerate(chunks, start=1)
]
return json.dumps(payload, ensure_ascii=False)
def _estimate_cost(model: ModelConfig | None, input_token: int, output_token: int) -> tuple[Decimal | None, str | None]:
if model is None:
return None, None
input_price = Decimal(model.input_price_per_1k or 0)
output_price = Decimal(model.output_price_per_1k or 0)
if input_price <= 0 and output_price <= 0:
return None, model.currency
value = (Decimal(input_token) / Decimal(1000) * input_price) + (Decimal(output_token) / Decimal(1000) * output_price)
return value.quantize(Decimal("0.000001"), rounding=ROUND_HALF_UP), model.currency
def _default_route_reason(retrieve_count: int) -> str:
return "当前默认正式模型;本轮命中知识库" if retrieve_count > 0 else "当前默认正式模型;本轮未命中知识库"
def _question_type(retrieve_count: int) -> str:
return "knowledge_grounded" if retrieve_count > 0 else "general_chat"

View File

@@ -154,6 +154,7 @@ class ChatService:
cost_ms=cost_ms,
error_message=str(exc),
retrieved_chunks=rag_result.chunks if rag_result is not None else None,
model_id=None,
)
db.commit()
raise HTTPException(
@@ -204,6 +205,7 @@ class ChatService:
output_token=completion.output_token,
cost_ms=cost_ms,
retrieved_chunks=rag_result.chunks,
model_id=completion.model_id,
)
db.commit()
return completion.answer

View File

@@ -117,6 +117,7 @@ class ChatStreamService:
cost_ms=cost_ms,
error_message=str(exc),
retrieved_chunks=rag_result.chunks if rag_result is not None else None,
model_id=model_response.model_id if model_response is not None else None,
)
_mark_retrieval_failed(db, rag_result, str(exc), cost_ms)
db.commit()
@@ -137,6 +138,7 @@ class ChatStreamService:
cost_ms=cost_ms,
error_message="模型未返回有效内容",
retrieved_chunks=rag_result.chunks if rag_result is not None else None,
model_id=model_response.model_id if model_response is not None else None,
)
_mark_retrieval_failed(db, rag_result, "模型未返回有效内容", cost_ms)
db.commit()
@@ -185,6 +187,7 @@ class ChatStreamService:
output_token=_rough_token_count(answer),
cost_ms=cost_ms,
retrieved_chunks=rag_result.chunks if rag_result is not None else None,
model_id=model_response.model_id if model_response is not None else None,
)
db.commit()
@@ -279,6 +282,7 @@ class ChatStreamService:
cost_ms=cost_ms,
error_message=str(exc),
retrieved_chunks=rag_result.chunks if rag_result is not None else None,
model_id=model_response.model_id if model_response is not None else None,
)
db.commit()
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=str(exc)) from exc
@@ -298,6 +302,7 @@ class ChatStreamService:
cost_ms=cost_ms,
error_message="模型未返回有效内容",
retrieved_chunks=rag_result.chunks if rag_result is not None else None,
model_id=model_response.model_id if model_response is not None else None,
)
db.commit()
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="模型未返回有效内容")
@@ -381,6 +386,7 @@ def _write_success(
output_token=_rough_token_count(answer),
cost_ms=cost_ms,
retrieved_chunks=rag_result.chunks if rag_result is not None else None,
model_id=model_response.model_id if model_response is not None else None,
)
if rag_result is not None and rag_result.retrieval_log_id:
retrieval_log = db.get(KnowledgeRetrievalLog, rag_result.retrieval_log_id)