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

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