feat: add ai cost tracking foundation
This commit is contained in:
@@ -231,7 +231,9 @@ def ai_logs(
|
||||
AiRequestLog.id, AiRequestLog.session_id, AiRequestLog.message_id, AiRequestLog.user_id,
|
||||
AiRequestLog.model_name, AiRequestLog.knowledge_ids, AiRequestLog.retrieve_count,
|
||||
AiRequestLog.input_token, AiRequestLog.output_token, AiRequestLog.total_token,
|
||||
AiRequestLog.cost_ms, AiRequestLog.status, AiRequestLog.error_message, AiRequestLog.created_at,
|
||||
AiRequestLog.cost_ms, AiRequestLog.estimated_cost, AiRequestLog.currency,
|
||||
AiRequestLog.route_reason, AiRequestLog.question_type, AiRequestLog.knowledge_hit,
|
||||
AiRequestLog.status, AiRequestLog.error_message, AiRequestLog.created_at,
|
||||
)).order_by(AiRequestLog.created_at.desc()).offset((page - 1) * pageSize).limit(pageSize)
|
||||
).all()
|
||||
return api_success(page_result([_ai_log_dict(item, include_chunks=False) for item in logs], total=total, page=page, page_size=pageSize))
|
||||
@@ -373,6 +375,11 @@ def _ai_log_dict(log: AiRequestLog, *, include_prompt: bool = False, include_chu
|
||||
"outputToken": log.output_token,
|
||||
"totalToken": log.total_token,
|
||||
"costMs": log.cost_ms,
|
||||
"estimatedCost": float(log.estimated_cost) if log.estimated_cost is not None else None,
|
||||
"currency": log.currency,
|
||||
"routeReason": log.route_reason,
|
||||
"questionType": log.question_type,
|
||||
"knowledgeHit": bool(log.knowledge_hit),
|
||||
"status": log.status,
|
||||
"errorMessage": log.error_message,
|
||||
"createdAt": log.created_at,
|
||||
|
||||
@@ -280,6 +280,14 @@ def create_model(
|
||||
extra_params=payload.extraParams,
|
||||
remark=payload.remark,
|
||||
timeout_second=payload.timeoutSecond,
|
||||
input_price_per_1k=payload.inputPricePer1k,
|
||||
output_price_per_1k=payload.outputPricePer1k,
|
||||
currency=payload.currency or "CNY",
|
||||
usage_scenarios=payload.usageScenarios,
|
||||
allow_summary=payload.allowSummary,
|
||||
allow_report=payload.allowReport,
|
||||
allow_fixed_info=payload.allowFixedInfo,
|
||||
allow_deep_chat=payload.allowDeepChat,
|
||||
enabled=0,
|
||||
)
|
||||
db.add(model)
|
||||
@@ -323,6 +331,14 @@ def update_model(
|
||||
model.extra_params = payload.extraParams
|
||||
model.remark = payload.remark
|
||||
model.timeout_second = payload.timeoutSecond
|
||||
model.input_price_per_1k = payload.inputPricePer1k
|
||||
model.output_price_per_1k = payload.outputPricePer1k
|
||||
model.currency = payload.currency or "CNY"
|
||||
model.usage_scenarios = payload.usageScenarios
|
||||
model.allow_summary = payload.allowSummary
|
||||
model.allow_report = payload.allowReport
|
||||
model.allow_fixed_info = payload.allowFixedInfo
|
||||
model.allow_deep_chat = payload.allowDeepChat
|
||||
db.add(model)
|
||||
OperationLogService.write(db, admin_id=current_admin.id, module="model", action="update", target_id=model.id)
|
||||
db.commit()
|
||||
@@ -433,6 +449,14 @@ def _model_dict(model: ModelConfig) -> dict:
|
||||
"remark": model.remark,
|
||||
"timeoutSecond": model.timeout_second,
|
||||
"enabled": model.enabled,
|
||||
"inputPricePer1k": float(model.input_price_per_1k) if model.input_price_per_1k is not None else None,
|
||||
"outputPricePer1k": float(model.output_price_per_1k) if model.output_price_per_1k is not None else None,
|
||||
"currency": model.currency,
|
||||
"usageScenarios": model.usage_scenarios,
|
||||
"allowSummary": model.allow_summary,
|
||||
"allowReport": model.allow_report,
|
||||
"allowFixedInfo": model.allow_fixed_info,
|
||||
"allowDeepChat": model.allow_deep_chat,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -49,6 +49,14 @@ class ModelConfig(Base):
|
||||
remark: Mapped[str | None] = mapped_column(String(255), nullable=True)
|
||||
timeout_second: Mapped[int] = mapped_column(Integer, default=30, nullable=False)
|
||||
enabled: Mapped[int] = mapped_column(default=0, nullable=False)
|
||||
input_price_per_1k: Mapped[Decimal | None] = mapped_column(Numeric(12, 6), nullable=True)
|
||||
output_price_per_1k: Mapped[Decimal | None] = mapped_column(Numeric(12, 6), nullable=True)
|
||||
currency: Mapped[str] = mapped_column(String(10), default="CNY", nullable=False)
|
||||
usage_scenarios: Mapped[str | None] = mapped_column(String(255), nullable=True)
|
||||
allow_summary: Mapped[int] = mapped_column(Integer, default=1, nullable=False)
|
||||
allow_report: Mapped[int] = mapped_column(Integer, default=1, nullable=False)
|
||||
allow_fixed_info: Mapped[int] = mapped_column(Integer, default=1, nullable=False)
|
||||
allow_deep_chat: Mapped[int] = mapped_column(Integer, default=1, nullable=False)
|
||||
|
||||
|
||||
class SystemConfig(Base):
|
||||
|
||||
@@ -2,7 +2,9 @@ from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import BigInteger, DateTime, Integer, String, Text, func
|
||||
from decimal import Decimal
|
||||
|
||||
from sqlalchemy import BigInteger, DateTime, Integer, Numeric, String, Text, func
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base
|
||||
@@ -17,6 +19,7 @@ class AiRequestLog(Base):
|
||||
session_id: Mapped[int | None] = mapped_column(BigInteger, index=True, nullable=True)
|
||||
message_id: Mapped[int | None] = mapped_column(BigInteger, index=True, nullable=True)
|
||||
user_id: Mapped[int | None] = mapped_column(BigInteger, index=True, nullable=True)
|
||||
model_id: Mapped[int | None] = mapped_column(BigInteger, index=True, nullable=True)
|
||||
model_name: Mapped[str | None] = mapped_column(String(100), nullable=True)
|
||||
prompt: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
knowledge_ids: Mapped[str | None] = mapped_column(String(500), nullable=True)
|
||||
@@ -26,6 +29,11 @@ class AiRequestLog(Base):
|
||||
output_token: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
total_token: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
cost_ms: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
estimated_cost: Mapped[Decimal | None] = mapped_column(Numeric(18, 6), nullable=True)
|
||||
currency: Mapped[str | None] = mapped_column(String(10), nullable=True)
|
||||
route_reason: Mapped[str | None] = mapped_column(String(255), nullable=True)
|
||||
question_type: Mapped[str | None] = mapped_column(String(50), nullable=True)
|
||||
knowledge_hit: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
status: Mapped[str] = mapped_column(String(20), nullable=False)
|
||||
error_message: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), nullable=False)
|
||||
|
||||
@@ -35,6 +35,8 @@ class DashboardStats(BaseModel):
|
||||
inputToken: int
|
||||
outputToken: int
|
||||
totalToken: int
|
||||
estimatedCost: float | None = None
|
||||
costCurrency: str | None = None
|
||||
|
||||
|
||||
class AdminUserUpdateRequest(BaseModel):
|
||||
@@ -161,6 +163,14 @@ class ModelSaveRequest(BaseModel):
|
||||
extraParams: str | None = None
|
||||
remark: str | None = Field(default=None, max_length=255)
|
||||
timeoutSecond: int = Field(default=30, ge=1, le=300)
|
||||
inputPricePer1k: float | None = Field(default=None, ge=0, le=1000000)
|
||||
outputPricePer1k: float | None = Field(default=None, ge=0, le=1000000)
|
||||
currency: str = Field(default="CNY", max_length=10)
|
||||
usageScenarios: str | None = Field(default=None, max_length=255)
|
||||
allowSummary: int = Field(default=1, ge=0, le=1)
|
||||
allowReport: int = Field(default=1, ge=0, le=1)
|
||||
allowFixedInfo: int = Field(default=1, ge=0, le=1)
|
||||
allowDeepChat: int = Field(default=1, ge=0, le=1)
|
||||
|
||||
|
||||
class EnableModelRequest(BaseModel):
|
||||
|
||||
@@ -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