feat: add ai cost tracking foundation
This commit is contained in:
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user