feat: route background ai workloads by model
This commit is contained in:
@@ -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))
|
||||
Reference in New Issue
Block a user