feat: route background ai workloads by model

This commit is contained in:
2026-07-31 17:25:16 +08:00
parent 0008903e8d
commit da313f88ed
22 changed files with 714 additions and 63 deletions

View File

@@ -0,0 +1,52 @@
"""add multi-model routing support
Revision ID: 0022_model_routing
Revises: 0021_periodic_report_async_jobs
"""
from alembic import op
import sqlalchemy as sa
revision = "0022_model_routing"
down_revision = "0021_periodic_report_async_jobs"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"sys_model",
sa.Column("is_default", sa.Integer(), nullable=False, server_default="0"),
)
op.create_index(
"ix_sys_model_available_default",
"sys_model",
["enabled", "is_default", "id"],
unique=False,
)
# 兼容旧数据:原来唯一启用的模型直接成为默认主模型。
op.execute(
"""
UPDATE sys_model
SET is_default = 1
WHERE enabled = 1
AND id = (
SELECT selected.id
FROM (
SELECT id
FROM sys_model
WHERE enabled = 1
ORDER BY id DESC
LIMIT 1
) AS selected
)
"""
)
def downgrade() -> None:
# 回滚到旧语义前只保留默认模型为启用,避免旧代码随机选到分流模型。
op.execute("UPDATE sys_model SET enabled = CASE WHEN is_default = 1 THEN 1 ELSE 0 END")
op.drop_index("ix_sys_model_available_default", table_name="sys_model")
op.drop_column("sys_model", "is_default")

View File

@@ -19,6 +19,7 @@ from app.models.knowledge import Knowledge
from app.schemas.admin import (
AgentDebugRequest,
AgentRuntimeConfigSaveRequest,
DefaultModelRequest,
EnableModelRequest,
ModelSaveRequest,
PromptSaveRequest,
@@ -246,7 +247,13 @@ def save_agent_runtime_config(
@router.get("/model/list")
def list_models(db: Session = Depends(get_db), current_admin: Admin = Depends(get_current_admin)) -> dict:
models = db.scalars(select(ModelConfig).order_by(ModelConfig.id.desc())).all()
models = db.scalars(
select(ModelConfig).order_by(
ModelConfig.is_default.desc(),
ModelConfig.enabled.desc(),
ModelConfig.id.desc(),
)
).all()
return api_success([_model_dict(model) for model in models])
@@ -289,6 +296,7 @@ def create_model(
allow_fixed_info=payload.allowFixedInfo,
allow_deep_chat=payload.allowDeepChat,
enabled=0,
is_default=0,
)
db.add(model)
db.flush()
@@ -351,14 +359,46 @@ def enable_model(
payload: EnableModelRequest,
db: Session = Depends(get_db),
current_admin: Admin = Depends(get_current_admin),
) -> dict:
target = db.get(ModelConfig, payload.modelId)
if target is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模型不存在")
if payload.enabled == 0 and target.is_default == 1:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="默认主模型不能直接停用,请先设置另一个默认主模型",
)
target.enabled = payload.enabled
if payload.enabled == 1 and _explicit_default_model(db) is None:
target.is_default = 1
db.add(target)
action = "enable" if payload.enabled == 1 else "disable"
OperationLogService.write(db, admin_id=current_admin.id, module="model", action=action, target_id=target.id)
db.commit()
return api_success()
@router.post("/model/default")
def set_default_model(
payload: DefaultModelRequest,
db: Session = Depends(get_db),
current_admin: Admin = Depends(get_current_admin),
) -> dict:
target = db.get(ModelConfig, payload.modelId)
if target is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模型不存在")
for model in db.scalars(select(ModelConfig)).all():
model.enabled = 1 if model.id == target.id else 0
model.is_default = 1 if model.id == target.id else 0
if model.id == target.id:
model.enabled = 1
db.add(model)
OperationLogService.write(db, admin_id=current_admin.id, module="model", action="enable", target_id=target.id)
OperationLogService.write(
db,
admin_id=current_admin.id,
module="model",
action="set_default",
target_id=target.id,
)
db.commit()
return api_success()
@@ -372,6 +412,11 @@ def delete_model(
model = db.get(ModelConfig, model_id)
if model is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模型不存在")
if model.is_default == 1:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="不能删除默认主模型,请先设置另一个默认主模型",
)
db.delete(model)
OperationLogService.write(db, admin_id=current_admin.id, module="model", action="delete", target_id=model.id)
db.commit()
@@ -449,6 +494,7 @@ def _model_dict(model: ModelConfig) -> dict:
"remark": model.remark,
"timeoutSecond": model.timeout_second,
"enabled": model.enabled,
"isDefault": model.is_default,
"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,
@@ -461,9 +507,13 @@ def _model_dict(model: ModelConfig) -> dict:
def _enabled_model(db: Session) -> ModelConfig | None:
return ModelClientService._get_enabled_model(db)
def _explicit_default_model(db: Session) -> ModelConfig | None:
return db.scalar(
select(ModelConfig)
.where(ModelConfig.enabled == 1)
.where(ModelConfig.enabled == 1, ModelConfig.is_default == 1)
.order_by(ModelConfig.id.desc())
.limit(1)
)

View File

@@ -49,6 +49,7 @@ 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)
is_default: 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)

View File

@@ -26,6 +26,17 @@ class AdminRead(ORMModel):
status: int
class CostBreakdownItem(BaseModel):
scene: str
modelName: str
requestCount: int
inputToken: int
outputToken: int
totalToken: int
estimatedCost: float
currency: str | None = None
class DashboardStats(BaseModel):
userCount: int
sessionCount: int
@@ -37,6 +48,7 @@ class DashboardStats(BaseModel):
totalToken: int
estimatedCost: float | None = None
costCurrency: str | None = None
costBreakdown: list[CostBreakdownItem] = Field(default_factory=list)
class AdminUserUpdateRequest(BaseModel):
@@ -175,6 +187,11 @@ class ModelSaveRequest(BaseModel):
class EnableModelRequest(BaseModel):
modelId: int
enabled: int = Field(default=1, ge=0, le=1)
class DefaultModelRequest(BaseModel):
modelId: int
class SystemConfigSaveRequest(BaseModel):

View File

@@ -3,7 +3,7 @@ from __future__ import annotations
from datetime import UTC, datetime
from fastapi import HTTPException, status
from sqlalchemy import func, select
from sqlalchemy import func, select, true
from sqlalchemy.orm import Session
from app.core.config import get_settings
@@ -133,23 +133,71 @@ class AdminDashboardService:
msg_filter.append(ChatMessage.created_at <= end)
ai_filter.append(AiRequestLog.created_at <= end)
user_where = and_(true(), *user_filter)
session_where = and_(true(), *session_filter)
message_where = and_(true(), *msg_filter)
ai_where = and_(true(), *ai_filter)
cost_currency = db.scalar(
select(AiRequestLog.currency)
.where(and_(*ai_filter), AiRequestLog.currency.is_not(None))
.where(ai_where, AiRequestLog.currency.is_not(None))
.order_by(AiRequestLog.id.desc())
.limit(1)
)
cost_breakdown = [
{
"scene": question_type or "unknown",
"modelName": model_name or "未知模型",
"requestCount": int(request_count or 0),
"inputToken": int(input_token or 0),
"outputToken": int(output_token or 0),
"totalToken": int(total_token or 0),
"estimatedCost": float(estimated_cost or 0),
"currency": currency,
}
for (
question_type,
model_name,
currency,
request_count,
input_token,
output_token,
total_token,
estimated_cost,
) in db.execute(
select(
AiRequestLog.question_type,
AiRequestLog.model_name,
AiRequestLog.currency,
func.count(AiRequestLog.id),
func.coalesce(func.sum(AiRequestLog.input_token), 0),
func.coalesce(func.sum(AiRequestLog.output_token), 0),
func.coalesce(func.sum(AiRequestLog.total_token), 0),
func.coalesce(func.sum(AiRequestLog.estimated_cost), 0),
)
.where(ai_where)
.group_by(
AiRequestLog.question_type,
AiRequestLog.model_name,
AiRequestLog.currency,
)
.order_by(func.count(AiRequestLog.id).desc())
).all()
]
currencies = {item["currency"] for item in cost_breakdown if item["currency"]}
if len(currencies) > 1:
cost_currency = "MIXED"
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,
"messageCount": db.scalar(select(func.count(ChatMessage.id)).where(and_(*msg_filter))) or 0,
"aiRequestCount": db.scalar(select(func.count(AiRequestLog.id)).where(and_(*ai_filter))) or 0,
"userCount": db.scalar(select(func.count(User.id)).where(user_where)) or 0,
"sessionCount": db.scalar(select(func.count(ChatSession.id)).where(session_where)) or 0,
"messageCount": db.scalar(select(func.count(ChatMessage.id)).where(message_where)) or 0,
"aiRequestCount": db.scalar(select(func.count(AiRequestLog.id)).where(ai_where)) or 0,
"knowledgeCount": db.scalar(select(func.count(Knowledge.id))) or 0,
"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),
"inputToken": db.scalar(select(func.coalesce(func.sum(AiRequestLog.input_token), 0)).where(ai_where)) or 0,
"outputToken": db.scalar(select(func.coalesce(func.sum(AiRequestLog.output_token), 0)).where(ai_where)) or 0,
"totalToken": db.scalar(select(func.coalesce(func.sum(AiRequestLog.total_token), 0)).where(ai_where)) or 0,
"estimatedCost": float(db.scalar(select(func.coalesce(func.sum(AiRequestLog.estimated_cost), 0)).where(ai_where)) or 0),
"costCurrency": cost_currency,
"costBreakdown": cost_breakdown,
}

View File

@@ -16,9 +16,9 @@ class AiRequestLogService:
def write_success(
db: Session,
*,
session_id: int,
session_id: int | None,
message_id: int | None,
user_id: int,
user_id: int | None,
model_name: str,
prompt: str,
knowledge_ids: str,
@@ -61,9 +61,9 @@ class AiRequestLogService:
def write_failed(
db: Session,
*,
session_id: int,
session_id: int | None,
message_id: int | None,
user_id: int,
user_id: int | None,
model_name: str | None,
prompt: str | None,
knowledge_ids: str | None,

View File

@@ -14,7 +14,7 @@ from app.models.growth import GrowthProfileRevision, TopicSummary, UserGrowthPro
from app.models.user import User
from app.services.entitlement_service import EntitlementService
from app.services.external_errors import ExternalServiceError
from app.services.model_service import ModelClientService
from app.services.tracked_generation_service import TrackedGenerationService
from app.services.topic_session_service import TopicSessionService
@@ -62,9 +62,16 @@ class GrowthProfileService:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="当前主题还没有可沉淀的对话内容")
conversation = _messages_text(messages)
prompt = _topic_summary_prompt(topic, conversation)
model_name = _enabled_model_name(db)
model_name = None
try:
raw = ModelClientService.generate_text_or_raise(db, prompt)
completion = TrackedGenerationService.generate(
db,
prompt=prompt,
scenario="summary",
user_id=user.id,
)
raw = completion.answer
model_name = completion.model_name
parsed = _parse_summary_json(raw)
data = parsed if parsed and "summary" in parsed else _fallback_summary(topic, conversation, raw)
status_value = "success"
@@ -96,7 +103,13 @@ class GrowthProfileService:
prompt = _growth_profile_prompt(profile, topic_summary)
try:
raw = ModelClientService.generate_text_or_raise(db, prompt)
completion = TrackedGenerationService.generate(
db,
prompt=prompt,
scenario="summary",
user_id=user.id,
)
raw = completion.answer
parsed = _parse_summary_json(raw)
data = parsed if parsed and "profileText" in parsed else _fallback_profile(profile, topic_summary, raw)
except ExternalServiceError:
@@ -326,10 +339,5 @@ def _limit(text: str, max_len: int) -> str:
return text[:max_len].strip()
def _enabled_model_name(db: Session) -> str | None:
model = ModelClientService._get_enabled_model(db)
return model.model_name if model is not None else None
def _now() -> datetime:
return datetime.now(UTC).replace(tzinfo=None)

View File

@@ -12,7 +12,6 @@ from sqlalchemy import or_, select
from sqlalchemy.orm import Session
from app.core.config import get_settings
from app.models.ai_config import ModelConfig
from app.services.chat_context_service import ChatContextService
from app.models.knowledge import (
Knowledge,
@@ -33,6 +32,7 @@ from app.services.knowledge_pipeline_service import (
from app.services.knowledge_catalog_cache_service import KnowledgeCatalogCacheService
from app.services.knowledge_service import KnowledgeScope
from app.services.model_service import _call_configured_model, _system_config_bool
from app.services.model_routing_service import ModelRoutingService
from app.services.rag_service import PromptService, RagResult, RetrievedChunk
SAFETY_RULE_VERSION = "minimum-safety-v1"
@@ -264,9 +264,7 @@ class KnowledgeAgentService:
"recentHistory": history_text,
}
fallback = cls._fallback_rewrite(question, recent)
model = db.scalar(
select(ModelConfig).where(ModelConfig.enabled == 1).order_by(ModelConfig.id.desc()).limit(1)
)
model = ModelRoutingService.default_model(db)
if model is None or _system_config_bool(db, "mock_model_enabled", get_settings().mock_model_enabled):
return fallback, cls._trace(
"rewrite_contextual_question",
@@ -444,7 +442,7 @@ class KnowledgeAgentService:
async def _rerank(cls, db: Session, question: str, candidates: list[Candidate], trace: list[dict], started: float) -> None:
if not candidates:
return
model = db.scalar(select(ModelConfig).where(ModelConfig.enabled == 1).order_by(ModelConfig.id.desc()).limit(1))
model = ModelRoutingService.default_model(db)
if model is None or _system_config_bool(db, "mock_model_enabled", get_settings().mock_model_enabled):
trace.append(cls._trace("Rerank", len(trace) + 1, {"candidateCount": len(candidates)}, {"count": len(candidates), "mode": "lexical_fallback", "candidates": [cls._candidate_trace(x) for x in candidates]}, started))
return

View File

@@ -0,0 +1,89 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Literal
from sqlalchemy import desc, select
from sqlalchemy.orm import Session
from app.models.ai_config import ModelConfig
ModelScenario = Literal["report", "summary", "fixed_info", "deep_chat"]
SCENARIO_LABELS: dict[ModelScenario, str] = {
"report": "周期报告",
"summary": "摘要沉淀",
"fixed_info": "固定信息",
"deep_chat": "深度对话",
}
_SCENARIO_FIELDS = {
"report": ModelConfig.allow_report,
"summary": ModelConfig.allow_summary,
"fixed_info": ModelConfig.allow_fixed_info,
"deep_chat": ModelConfig.allow_deep_chat,
}
@dataclass(frozen=True)
class ModelRoute:
model: ModelConfig | None
scenario: ModelScenario
reason: str
fallback_used: bool
class ModelRoutingService:
"""Centralizes deterministic model selection without changing live-chat routing."""
@staticmethod
def default_model(db: Session) -> ModelConfig | None:
model = db.scalar(
select(ModelConfig)
.where(ModelConfig.enabled == 1, ModelConfig.is_default == 1)
.order_by(ModelConfig.id.desc())
.limit(1)
)
if model is not None:
return model
# Migration/partial rollout safety: old databases may have an available
# model before one has been explicitly marked as default.
return db.scalar(
select(ModelConfig)
.where(ModelConfig.enabled == 1)
.order_by(ModelConfig.id.desc())
.limit(1)
)
@classmethod
def resolve(cls, db: Session, scenario: ModelScenario) -> ModelRoute:
default = cls.default_model(db)
capability = _SCENARIO_FIELDS[scenario]
candidate = db.scalar(
select(ModelConfig)
.where(ModelConfig.enabled == 1, capability == 1)
.order_by(desc(ModelConfig.is_default), ModelConfig.id.desc())
.limit(1)
)
label = SCENARIO_LABELS[scenario]
if candidate is not None:
source = "默认主模型" if candidate.is_default == 1 else "场景可用模型"
return ModelRoute(
model=candidate,
scenario=scenario,
reason=f"场景分流:{label};选择{source}",
fallback_used=False,
)
if default is not None:
return ModelRoute(
model=default,
scenario=scenario,
reason=f"场景分流:{label}无匹配模型;回退默认主模型",
fallback_used=True,
)
return ModelRoute(
model=None,
scenario=scenario,
reason=f"场景分流:{label}无可用模型",
fallback_used=True,
)

View File

@@ -13,6 +13,7 @@ from sqlalchemy.orm import Session
from app.core.config import get_settings
from app.models.ai_config import ModelConfig, SystemConfig
from app.services.external_errors import ExternalServiceError
from app.services.model_routing_service import ModelRoutingService, ModelScenario
from app.services.rag_service import NO_HIT_ANSWER, RagResult
from app.services.secret_service import SecretService
@@ -24,6 +25,7 @@ class ModelCompletion:
model_name: str
input_token: int
output_token: int
route_reason: str | None = None
class ModelClientService:
@@ -83,6 +85,41 @@ class ModelClientService:
rag_result = RagResult(question=prompt, knowledge_scopes=[], chunks=[], prompt=prompt, allow_general_knowledge=True)
return _call_configured_model(model, rag_result, allow_no_hit=True)
@staticmethod
def generate_text_for_scenario(
db: Session,
prompt: str,
scenario: ModelScenario,
) -> ModelCompletion:
route = ModelRoutingService.resolve(db, scenario)
model = route.model
if _system_config_bool(db, "mock_model_enabled", get_settings().mock_model_enabled):
answer = prompt.strip()[-1200:]
model_name = model.model_name if model is not None else "mock-model"
else:
if model is None:
raise ExternalServiceError(
f"未启用可用于{scenario}场景的模型",
provider="model",
)
rag_result = RagResult(
question=prompt,
knowledge_scopes=[],
chunks=[],
prompt=prompt,
allow_general_knowledge=True,
)
answer = _call_configured_model(model, rag_result, allow_no_hit=True)
model_name = model.model_name
return ModelCompletion(
answer=answer,
model_id=model.id if model is not None else None,
model_name=model_name,
input_token=_rough_token_count(prompt),
output_token=_rough_token_count(answer),
route_reason=route.reason,
)
@staticmethod
async def summarize_or_raise_async(
db: Session,
@@ -99,12 +136,7 @@ class ModelClientService:
@staticmethod
def _get_enabled_model(db: Session) -> ModelConfig | None:
return db.scalar(
select(ModelConfig)
.where(ModelConfig.enabled == 1)
.order_by(ModelConfig.id.desc())
.limit(1)
)
return ModelRoutingService.default_model(db)
@staticmethod
def test_model(model: ModelConfig) -> dict[str, Any]:

View File

@@ -7,7 +7,6 @@ from dataclasses import dataclass
from typing import Any
import httpx
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.core.config import get_settings
@@ -30,6 +29,7 @@ from app.services.model_service import (
_system_and_turn_messages,
_system_config_bool,
)
from app.services.model_routing_service import ModelRoutingService
from app.services.rag_service import RagResult
@@ -120,12 +120,7 @@ class ModelStreamService:
def _get_enabled_model(db: Session) -> ModelConfig | None:
return db.scalar(
select(ModelConfig)
.where(ModelConfig.enabled == 1)
.order_by(ModelConfig.id.desc())
.limit(1)
)
return ModelRoutingService.default_model(db)
def _stream_configured_model(model: ModelConfig, rag_result: RagResult) -> Iterator[str]:

View File

@@ -12,7 +12,7 @@ from sqlalchemy.orm import Session
from app.core.config import get_settings
from app.models.growth import PeriodicReport, TopicSummary, UserGrowthProfile
from app.models.user import User
from app.services.model_service import ModelClientService
from app.services.tracked_generation_service import TrackedGenerationService
ReportType = Literal["weekly", "monthly", "stage"]
@@ -177,9 +177,14 @@ class PeriodicReportService:
try:
prompt = _report_prompt(user=user, report_type=report_type, period_start=period_start, period_end=period_end, summaries=summaries, profile=profile)
report.content = ModelClientService.generate_text_or_raise(db, prompt).strip() or _fallback_report(summaries)
model = ModelClientService._get_enabled_model(db)
report.model_name = model.model_name if model else None
completion = TrackedGenerationService.generate(
db,
prompt=prompt,
scenario="report",
user_id=user.id,
)
report.content = completion.answer.strip() or _fallback_report(summaries)
report.model_name = completion.model_name
report.status = "success"
report.error_message = None
except Exception as exc:

View File

@@ -0,0 +1,64 @@
from __future__ import annotations
from time import perf_counter
from sqlalchemy.orm import Session
from app.services.ai_request_log_service import AiRequestLogService
from app.services.model_routing_service import ModelRoutingService, ModelScenario
from app.services.model_service import ModelClientService, ModelCompletion
class TrackedGenerationService:
"""Runs non-chat generation and records its actual model, tokens and cost."""
@staticmethod
def generate(
db: Session,
*,
prompt: str,
scenario: ModelScenario,
user_id: int | None,
) -> ModelCompletion:
started = perf_counter()
try:
completion = ModelClientService.generate_text_for_scenario(db, prompt, scenario)
AiRequestLogService.write_success(
db,
session_id=None,
message_id=None,
user_id=user_id,
model_id=completion.model_id,
model_name=completion.model_name,
prompt=prompt,
knowledge_ids="",
retrieve_count=0,
input_token=completion.input_token,
output_token=completion.output_token,
cost_ms=_elapsed_ms(started),
route_reason=completion.route_reason,
question_type=f"background_{scenario}",
)
return completion
except Exception as exc:
route = ModelRoutingService.resolve(db, scenario)
AiRequestLogService.write_failed(
db,
session_id=None,
message_id=None,
user_id=user_id,
model_id=route.model.id if route.model is not None else None,
model_name=route.model.model_name if route.model is not None else None,
prompt=prompt,
knowledge_ids="",
retrieve_count=0,
cost_ms=_elapsed_ms(started),
route_reason=route.reason,
question_type=f"background_{scenario}",
error_message=str(exc)[:2000],
)
raise
def _elapsed_ms(started: float) -> int:
return max(0, int((perf_counter() - started) * 1000))

View File

@@ -9,6 +9,7 @@ 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.admin_service import AdminDashboardService
from app.services.ai_request_log_service import AiRequestLogService
@@ -58,3 +59,52 @@ def test_ai_request_log_estimates_cost_from_model_price():
assert log.question_type == "knowledge_grounded"
assert "命中知识库" in (log.route_reason or "")
def test_dashboard_breaks_cost_down_by_actual_scene_and_model():
with _db() as db:
model = ModelConfig(
id=1,
provider="test",
api_type="openai_compatible",
model_name="report-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=None,
message_id=None,
user_id=3,
model_id=1,
model_name="report-model",
prompt="weekly report",
knowledge_ids="",
retrieve_count=0,
input_token=1000,
output_token=500,
cost_ms=120,
question_type="background_report",
route_reason="场景分流:周期报告",
)
db.commit()
stats = AdminDashboardService.stats(db)
assert stats["costBreakdown"] == [
{
"scene": "background_report",
"modelName": "report-model",
"requestCount": 1,
"inputToken": 1000,
"outputToken": 500,
"totalToken": 1500,
"estimatedCost": 0.005,
"currency": "CNY",
}
]

View File

@@ -0,0 +1,154 @@
from __future__ import annotations
from decimal import Decimal
import pytest
from fastapi import HTTPException
from sqlalchemy import create_engine
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool
from app.models import Base
from app.api.admin_settings import enable_model, set_default_model
from app.models.admin import Admin
from app.models.ai_config import ModelConfig, SystemConfig
from app.models.logs import AiRequestLog
from app.schemas.admin import DefaultModelRequest, EnableModelRequest
from app.services.model_routing_service import ModelRoutingService
from app.services.model_service import ModelClientService
from app.services.tracked_generation_service import TrackedGenerationService
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 _model(
model_id: int,
name: str,
*,
is_default: int,
allow_report: int,
allow_summary: int,
) -> ModelConfig:
return ModelConfig(
id=model_id,
provider="test",
api_type="openai_compatible",
model_name=name,
api_url="https://example.com",
api_key="secret",
enabled=1,
is_default=is_default,
allow_report=allow_report,
allow_summary=allow_summary,
allow_fixed_info=1,
allow_deep_chat=1,
input_price_per_1k=Decimal("0.002"),
output_price_per_1k=Decimal("0.006"),
currency="CNY",
timeout_second=30,
)
def test_background_scenario_routes_without_changing_live_default_model():
with _db() as db:
main = _model(1, "main-model", is_default=1, allow_report=0, allow_summary=1)
report = _model(2, "report-model", is_default=0, allow_report=1, allow_summary=0)
db.add_all([main, report])
db.commit()
report_route = ModelRoutingService.resolve(db, "report")
summary_route = ModelRoutingService.resolve(db, "summary")
assert report_route.model is report
assert report_route.fallback_used is False
assert summary_route.model is main
assert ModelClientService._get_enabled_model(db) is main
def test_background_scenario_falls_back_to_default_model():
with _db() as db:
main = _model(1, "main-model", is_default=1, allow_report=0, allow_summary=0)
db.add(main)
db.commit()
route = ModelRoutingService.resolve(db, "report")
assert route.model is main
assert route.fallback_used is True
assert "回退默认主模型" in route.reason
def test_tracked_background_generation_records_actual_route_tokens_and_cost():
with _db() as db:
db.add(SystemConfig(config_key="mock_model_enabled", config_value="true"))
report = _model(2, "report-model", is_default=0, allow_report=1, allow_summary=0)
db.add(report)
db.commit()
completion = TrackedGenerationService.generate(
db,
prompt="生成本周报告",
scenario="report",
user_id=7,
)
db.commit()
log = db.query(AiRequestLog).one()
assert completion.model_name == "report-model"
assert log.model_id == report.id
assert log.user_id == 7
assert log.question_type == "background_report"
assert log.total_token == completion.input_token + completion.output_token
assert log.estimated_cost is not None
assert "周期报告" in (log.route_reason or "")
def test_model_pool_keeps_one_default_and_rejects_disabling_last_default():
with _db() as db:
admin = Admin(id=1, username="admin", password="hash", name="管理员", status=1)
main = _model(1, "main-model", is_default=1, allow_report=1, allow_summary=1)
alternate = _model(2, "alternate-model", is_default=0, allow_report=1, allow_summary=1)
alternate.enabled = 0
db.add_all([admin, main, alternate])
db.commit()
enable_model(
EnableModelRequest(modelId=alternate.id, enabled=1),
db=db,
current_admin=admin,
)
db.refresh(main)
db.refresh(alternate)
assert main.is_default == 1
assert alternate.enabled == 1
set_default_model(
DefaultModelRequest(modelId=alternate.id),
db=db,
current_admin=admin,
)
db.refresh(main)
db.refresh(alternate)
assert main.is_default == 0
assert alternate.is_default == 1
enable_model(
EnableModelRequest(modelId=main.id, enabled=0),
db=db,
current_admin=admin,
)
with pytest.raises(HTTPException) as exc_info:
enable_model(
EnableModelRequest(modelId=alternate.id, enabled=0),
db=db,
current_admin=admin,
)
assert exc_info.value.status_code == 400

View File

@@ -12,9 +12,9 @@ from app.models.chat import ChatSession, TopicSession
from app.models.entitlement import EntitlementPlan, UserEntitlement
from app.models.growth import PeriodicReport, TopicSummary, UserGrowthProfile
from app.models.user import User
from app.services.model_service import ModelClientService
from app.services.periodic_report_service import PeriodicReportService, periodic_report_dict
from app.services.periodic_report_worker import PeriodicReportWorker, scheduled_period
from app.services.tracked_generation_service import TrackedGenerationService
def _db() -> Session:
@@ -150,9 +150,9 @@ def test_failed_async_report_is_retried(monkeypatch):
)
db.commit()
monkeypatch.setattr(
ModelClientService,
"generate_text_or_raise",
staticmethod(lambda _db, _prompt: (_ for _ in ()).throw(RuntimeError("模型暂时不可用"))),
TrackedGenerationService,
"generate",
staticmethod(lambda *_args, **_kwargs: (_ for _ in ()).throw(RuntimeError("模型暂时不可用"))),
)
report_id = PeriodicReportWorker.claim_next(db, worker_id="retry-worker", now=_now() + timedelta(seconds=1))