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

@@ -0,0 +1,66 @@
"""add ai cost tracking fields
Revision ID: 0018_ai_cost_tracking
Revises: 0017_share_drafts
"""
from alembic import op
import sqlalchemy as sa
revision = "0018_ai_cost_tracking"
down_revision = "0017_share_drafts"
branch_labels = None
depends_on = None
def upgrade() -> None:
inspector = sa.inspect(op.get_bind())
model_columns = {column["name"] for column in inspector.get_columns("sys_model")}
_add_column_if_missing(model_columns, "sys_model", sa.Column("input_price_per_1k", sa.Numeric(12, 6), nullable=True))
_add_column_if_missing(model_columns, "sys_model", sa.Column("output_price_per_1k", sa.Numeric(12, 6), nullable=True))
_add_column_if_missing(model_columns, "sys_model", sa.Column("currency", sa.String(10), nullable=False, server_default="CNY"))
_add_column_if_missing(model_columns, "sys_model", sa.Column("usage_scenarios", sa.String(255), nullable=True))
_add_column_if_missing(model_columns, "sys_model", sa.Column("allow_summary", sa.Integer(), nullable=False, server_default="1"))
_add_column_if_missing(model_columns, "sys_model", sa.Column("allow_report", sa.Integer(), nullable=False, server_default="1"))
_add_column_if_missing(model_columns, "sys_model", sa.Column("allow_fixed_info", sa.Integer(), nullable=False, server_default="1"))
_add_column_if_missing(model_columns, "sys_model", sa.Column("allow_deep_chat", sa.Integer(), nullable=False, server_default="1"))
log_columns = {column["name"] for column in inspector.get_columns("sys_ai_request_log")}
_add_column_if_missing(log_columns, "sys_ai_request_log", sa.Column("model_id", sa.BigInteger(), nullable=True))
_add_column_if_missing(log_columns, "sys_ai_request_log", sa.Column("estimated_cost", sa.Numeric(18, 6), nullable=True))
_add_column_if_missing(log_columns, "sys_ai_request_log", sa.Column("currency", sa.String(10), nullable=True))
_add_column_if_missing(log_columns, "sys_ai_request_log", sa.Column("route_reason", sa.String(255), nullable=True))
_add_column_if_missing(log_columns, "sys_ai_request_log", sa.Column("question_type", sa.String(50), nullable=True))
_add_column_if_missing(log_columns, "sys_ai_request_log", sa.Column("knowledge_hit", sa.Integer(), nullable=False, server_default="0"))
if "model_id" not in log_columns:
op.create_index("ix_sys_ai_request_log_model_id", "sys_ai_request_log", ["model_id"])
def downgrade() -> None:
inspector = sa.inspect(op.get_bind())
log_columns = {column["name"] for column in inspector.get_columns("sys_ai_request_log")}
if "model_id" in log_columns:
op.drop_index("ix_sys_ai_request_log_model_id", table_name="sys_ai_request_log")
for column in ["knowledge_hit", "question_type", "route_reason", "currency", "estimated_cost", "model_id"]:
if column in log_columns:
op.drop_column("sys_ai_request_log", column)
model_columns = {column["name"] for column in inspector.get_columns("sys_model")}
for column in [
"allow_deep_chat",
"allow_fixed_info",
"allow_report",
"allow_summary",
"usage_scenarios",
"currency",
"output_price_per_1k",
"input_price_per_1k",
]:
if column in model_columns:
op.drop_column("sys_model", column)
def _add_column_if_missing(columns: set[str], table: str, column: sa.Column) -> None:
if column.name not in columns:
op.add_column(table, column)

View File

@@ -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,

View File

@@ -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,
}

View File

@@ -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):

View File

@@ -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)

View File

@@ -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):

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)

View File

@@ -0,0 +1,60 @@
from __future__ import annotations
from decimal import Decimal
from sqlalchemy import create_engine
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool
from app.models import Base
from app.models.ai_config import ModelConfig
from app.models.logs import AiRequestLog
from app.services.ai_request_log_service import AiRequestLogService
def _db() -> Session:
engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool)
Base.metadata.create_all(engine)
return Session(engine)
def test_ai_request_log_estimates_cost_from_model_price():
with _db() as db:
model = ModelConfig(
id=1,
provider="test",
api_type="openai_compatible",
model_name="test-model",
api_url="https://example.com",
api_key="secret",
input_price_per_1k=Decimal("0.002"),
output_price_per_1k=Decimal("0.006"),
currency="CNY",
timeout_second=30,
)
db.add(model)
db.commit()
AiRequestLogService.write_success(
db,
session_id=1,
message_id=2,
user_id=3,
model_id=1,
model_name="test-model",
prompt="hello",
knowledge_ids="1",
retrieve_count=2,
input_token=1000,
output_token=500,
cost_ms=120,
)
db.commit()
log = db.query(AiRequestLog).one()
assert log.estimated_cost == Decimal("0.005000")
assert log.currency == "CNY"
assert log.knowledge_hit == 1
assert log.question_type == "knowledge_grounded"
assert "命中知识库" in (log.route_reason or "")