feat: route background ai workloads by model
This commit is contained in:
@@ -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")
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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]:
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
@@ -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",
|
||||
}
|
||||
]
|
||||
|
||||
154
ai_knowledge_base_v2/apps/backend/tests/test_model_routing.py
Normal file
154
ai_knowledge_base_v2/apps/backend/tests/test_model_routing.py
Normal 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
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user