feat: route background ai workloads by model
This commit is contained in:
@@ -305,7 +305,15 @@ async function loadDashboard() {
|
||||
|
||||
async function refreshStorage() { storageLoading.value = true; try { storageStats.value = await api.storageStats(true); ElMessage.success("存储统计已刷新"); } finally { storageLoading.value = false; } }
|
||||
function formatBytes(value: unknown) { if (typeof value !== "number") return "无法统计"; const units = ["B", "KB", "MB", "GB", "TB"]; let size = value, i = 0; while (size >= 1024 && i < units.length - 1) { size /= 1024; i += 1; } return `${size.toFixed(i ? 2 : 0)} ${units[i]}`; }
|
||||
function formatMoney(value?: number | null, currency = "CNY") { if (typeof value !== "number") return "-"; return `${currency || "CNY"} ${value.toFixed(6)}`; }
|
||||
function formatMoney(value?: number | null, currency = "CNY") { if (typeof value !== "number") return "-"; if (currency === "MIXED") return "多币种,见明细"; return `${currency || "CNY"} ${value.toFixed(6)}`; }
|
||||
function modelSceneLabel(scene: string) {
|
||||
return ({
|
||||
background_report: "周期报告",
|
||||
background_summary: "摘要沉淀",
|
||||
knowledge_grounded: "知识问答",
|
||||
general_chat: "通用对话",
|
||||
} as Record<string, string>)[scene] || scene || "未分类";
|
||||
}
|
||||
async function cleanupRetrievalLogs() { if (!cleanupBefore.value) return ElMessage.warning("请选择清理日期"); const estimate = await api.estimateRetrievalCleanup(cleanupBefore.value); await ElMessageBox.confirm(`预计清理 ${estimate.estimatedCount} 条检索日志及关联明细,操作日志不会被删除。`, "确认清理", { type: "warning" }); const result = await api.cleanupRetrievalLogs(cleanupBefore.value); ElMessage.success(`已清理 ${result.deleted} 条日志`); await loadCurrentMenu(); }
|
||||
async function saveRetention() { await api.saveRetrievalRetention(retentionDays.value); ElMessage.success(retentionDays.value == null ? "已设为永久保留" : `已设为保留 ${retentionDays.value} 天`); await loadCurrentMenu(); }
|
||||
|
||||
@@ -711,12 +719,25 @@ function applyQuickProvider(provider: string) {
|
||||
quickModelForm.modelName = target.models[0]?.value ?? "";
|
||||
}
|
||||
|
||||
async function enableModel(id: number) {
|
||||
await api.enableModel(id);
|
||||
ElMessage.success("模型已启用");
|
||||
async function setModelAvailability(row: ModelItem, enabled: number) {
|
||||
await api.setModelAvailability(row.id, enabled);
|
||||
ElMessage.success(enabled === 1 ? "模型已加入可用池" : "模型已停用");
|
||||
await loadCurrentMenu();
|
||||
}
|
||||
|
||||
async function setDefaultModel(row: ModelItem) {
|
||||
await api.setDefaultModel(row.id);
|
||||
ElMessage.success("默认主模型已更新");
|
||||
await loadCurrentMenu();
|
||||
}
|
||||
|
||||
async function handleModelCommand(command: string, row: ModelItem) {
|
||||
if (command === "enable") return setModelAvailability(row, 1);
|
||||
if (command === "disable") return setModelAvailability(row, 0);
|
||||
if (command === "default") return setDefaultModel(row);
|
||||
return deleteModel(row.id);
|
||||
}
|
||||
|
||||
async function deleteModel(id: number) {
|
||||
try {
|
||||
await ElMessageBox.confirm("确认删除该模型?删除后不可恢复。", "删除模型", {
|
||||
@@ -1108,6 +1129,23 @@ function formatRecordDateTime(value: string, boundary: "start" | "end") {
|
||||
<div class="stat"><span>总 Token</span><strong>{{ stats?.totalToken ?? 0 }}</strong></div>
|
||||
<div class="stat"><span>估算成本</span><strong>{{ formatMoney(stats?.estimatedCost, stats?.costCurrency || 'CNY') }}</strong></div>
|
||||
</div>
|
||||
<section class="cost-breakdown-panel">
|
||||
<div class="storage-panel-head">
|
||||
<div><h3>模型使用与成本</h3><p>按实际调用场景和模型统计;报告、摘要分流是否生效可在这里核对。</p></div>
|
||||
</div>
|
||||
<el-table :data="stats?.costBreakdown || []" empty-text="当前日期范围内暂无模型调用">
|
||||
<el-table-column label="调用场景" min-width="130">
|
||||
<template #default="{ row }">{{ modelSceneLabel(row.scene) }}</template>
|
||||
</el-table-column>
|
||||
<el-table-column prop="modelName" label="实际模型" min-width="180" show-overflow-tooltip />
|
||||
<el-table-column prop="requestCount" label="请求数" width="100" />
|
||||
<el-table-column prop="inputToken" label="输入 Token" width="130" />
|
||||
<el-table-column prop="outputToken" label="输出 Token" width="130" />
|
||||
<el-table-column label="估算成本" width="170">
|
||||
<template #default="{ row }">{{ formatMoney(row.estimatedCost, row.currency || 'CNY') }}</template>
|
||||
</el-table-column>
|
||||
</el-table>
|
||||
</section>
|
||||
<section class="storage-panel" v-loading="storageLoading"><div class="storage-panel-head"><div><h3>项目存储占用</h3><p>数据库、Redis 与附件分别统计;无权限或不支持时显示“无法统计”。</p></div><el-button :loading="storageLoading" @click="refreshStorage">手动刷新</el-button></div><div class="storage-grid"><div><span>项目总量</span><strong>{{ formatBytes(storageStats?.totalBytes) }}</strong></div><div><span>数据库</span><strong>{{ formatBytes(storageStats?.detail?.database?.bytes) }}</strong></div><div><span>Redis</span><strong>{{ formatBytes(storageStats?.detail?.redis?.bytes) }}</strong></div><div><span>附件/文件</span><strong>{{ formatBytes(storageStats?.detail?.files?.bytes) }}</strong></div><div><span>近7天增长</span><strong>{{ formatBytes(storageStats?.growth7DaysBytes) }}</strong></div><div><span>近30天增长</span><strong>{{ formatBytes(storageStats?.growth30DaysBytes) }}</strong></div></div><small>最后统计:{{ storageStats?.createdAt || '尚未统计' }}</small></section>
|
||||
</template>
|
||||
|
||||
@@ -1322,7 +1360,7 @@ function formatRecordDateTime(value: string, boundary: "start" | "end") {
|
||||
|
||||
<template v-if="activeMenu === 'models'">
|
||||
<div class="page-head inline">
|
||||
<div><h2>模型管理</h2><p>DeepSeek、MiniMax 可快速添加;其他供应商按 OpenAI 兼容或 Anthropic Messages 协议添加。</p></div>
|
||||
<div><h2>模型管理</h2><p>可同时启用多个模型,但只保留一个默认主模型;正式聊天始终使用主模型,报告和摘要按能力标签分流并自动回退主模型。</p></div>
|
||||
<el-button @click="resetModelForm">清空表单</el-button>
|
||||
</div>
|
||||
|
||||
@@ -1444,13 +1482,27 @@ function formatRecordDateTime(value: string, boundary: "start" | "end") {
|
||||
<el-table-column label="估算单价" width="180">
|
||||
<template #default="{ row }">{{ row.currency || 'CNY' }} {{ row.inputPricePer1k ?? '-' }}/{{ row.outputPricePer1k ?? '-' }}</template>
|
||||
</el-table-column>
|
||||
<el-table-column prop="usageScenarios" label="适用场景" min-width="160" show-overflow-tooltip />
|
||||
<el-table-column label="分流能力" min-width="220">
|
||||
<template #default="{ row }">
|
||||
<div class="model-capability-tags">
|
||||
<el-tag v-if="row.allowDeepChat === 1" size="small">深度对话</el-tag>
|
||||
<el-tag v-if="row.allowFixedInfo === 1" size="small">固定信息</el-tag>
|
||||
<el-tag v-if="row.allowSummary === 1" size="small" type="success">摘要</el-tag>
|
||||
<el-tag v-if="row.allowReport === 1" size="small" type="warning">报告</el-tag>
|
||||
</div>
|
||||
</template>
|
||||
</el-table-column>
|
||||
<el-table-column prop="baseUrl" label="Base URL" min-width="220" show-overflow-tooltip />
|
||||
<el-table-column prop="authType" label="鉴权" width="110" />
|
||||
<el-table-column prop="enabled" label="启用" width="90" />
|
||||
<el-table-column label="运行状态" width="150">
|
||||
<template #default="{ row }">
|
||||
<el-tag :type="row.enabled === 1 ? 'success' : 'info'">{{ row.enabled === 1 ? '可用' : '停用' }}</el-tag>
|
||||
<el-tag v-if="row.isDefault === 1" class="model-default-tag">主模型</el-tag>
|
||||
</template>
|
||||
</el-table-column>
|
||||
<el-table-column label="操作" width="190" fixed="right" align="center">
|
||||
<template #default="{ row }">
|
||||
<TableRowActions><el-button size="small" @click="editModel(row)">编辑</el-button><el-button size="small" @click="testModel(row.id)">测试</el-button><el-dropdown trigger="click" teleported @command="(command: string) => command === 'enable' ? enableModel(row.id) : deleteModel(row.id)"><el-button size="small">更多</el-button><template #dropdown><el-dropdown-menu><el-dropdown-item command="enable">启用</el-dropdown-item><el-dropdown-item command="delete" divided>删除</el-dropdown-item></el-dropdown-menu></template></el-dropdown></TableRowActions>
|
||||
<TableRowActions><el-button size="small" @click="editModel(row)">编辑</el-button><el-button size="small" @click="testModel(row.id)">测试</el-button><el-dropdown trigger="click" teleported @command="(command: string) => handleModelCommand(command, row)"><el-button size="small">更多</el-button><template #dropdown><el-dropdown-menu><el-dropdown-item v-if="row.enabled !== 1" command="enable">加入可用池</el-dropdown-item><el-dropdown-item v-if="row.enabled === 1 && row.isDefault !== 1" command="default">设为主模型</el-dropdown-item><el-dropdown-item v-if="row.enabled === 1 && row.isDefault !== 1" command="disable">停用</el-dropdown-item><el-dropdown-item v-if="row.isDefault !== 1" command="delete" divided>删除</el-dropdown-item></el-dropdown-menu></template></el-dropdown></TableRowActions>
|
||||
</template>
|
||||
</el-table-column>
|
||||
</el-table>
|
||||
|
||||
@@ -207,8 +207,10 @@ export const api = {
|
||||
request<ModelItem>("/admin/model", { method: "POST", body: JSON.stringify(payload) }),
|
||||
updateModel: (id: number, payload: Record<string, unknown>) =>
|
||||
request<ModelItem>(`/admin/model/${id}`, { method: "PUT", body: JSON.stringify(payload) }),
|
||||
enableModel: (modelId: number) =>
|
||||
request<null>("/admin/model/enable", { method: "POST", body: JSON.stringify({ modelId }) }),
|
||||
setModelAvailability: (modelId: number, enabled: number) =>
|
||||
request<null>("/admin/model/enable", { method: "POST", body: JSON.stringify({ modelId, enabled }) }),
|
||||
setDefaultModel: (modelId: number) =>
|
||||
request<null>("/admin/model/default", { method: "POST", body: JSON.stringify({ modelId }) }),
|
||||
deleteModel: (modelId: number) =>
|
||||
request<null>(`/admin/model/${modelId}`, { method: "DELETE" }),
|
||||
testModel: (modelId: number) =>
|
||||
|
||||
@@ -525,12 +525,15 @@ textarea {
|
||||
gap: 14px;
|
||||
}
|
||||
.storage-panel { margin-top: 18px; padding: 18px; border: 1px solid #dfe8e5; border-radius: 8px; background: #fff; }
|
||||
.cost-breakdown-panel { margin-top: 18px; padding: 18px; border: 1px solid #dfe8e5; border-radius: 8px; background: #fff; }
|
||||
.storage-panel-head { display: flex; justify-content: space-between; align-items: start; gap: 16px; }
|
||||
.storage-panel-head h3, .storage-panel-head p { margin: 0 0 6px; }
|
||||
.storage-grid { display: grid; grid-template-columns: repeat(4, minmax(0, 1fr)); gap: 12px; margin: 14px 0; }
|
||||
.storage-grid > div { padding: 14px; background: #f6f8f7; border-radius: 8px; }
|
||||
.storage-grid span, .storage-grid strong { display: block; }
|
||||
.storage-grid strong { margin-top: 8px; font-size: 20px; }
|
||||
.model-capability-tags { display: flex; flex-wrap: wrap; gap: 6px; }
|
||||
.model-default-tag { margin-left: 6px; }
|
||||
.cleanup-controls { display: flex; align-items: center; gap: 8px; white-space: nowrap; }
|
||||
.retention-panel { display: grid; grid-template-columns: minmax(260px, 1fr) 180px auto auto; align-items: center; gap: 12px; padding: 14px 16px; margin-bottom: 14px; border: 1px solid #dfe8e5; border-radius: 8px; background: #fff; }
|
||||
.retention-panel div span { display: block; margin-top: 4px; color: #667a73; font-size: 12px; }
|
||||
|
||||
@@ -29,6 +29,16 @@ export interface DashboardStats {
|
||||
totalToken: number;
|
||||
estimatedCost?: number | null;
|
||||
costCurrency?: string | null;
|
||||
costBreakdown: Array<{
|
||||
scene: string;
|
||||
modelName: string;
|
||||
requestCount: number;
|
||||
inputToken: number;
|
||||
outputToken: number;
|
||||
totalToken: number;
|
||||
estimatedCost: number;
|
||||
currency?: string | null;
|
||||
}>;
|
||||
}
|
||||
|
||||
export interface PromptDetail {
|
||||
@@ -195,6 +205,7 @@ export interface ModelItem {
|
||||
remark?: string | null;
|
||||
timeoutSecond: number;
|
||||
enabled: number;
|
||||
isDefault: number;
|
||||
inputPricePer1k?: number | null;
|
||||
outputPricePer1k?: number | null;
|
||||
currency: string;
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -163,6 +163,25 @@ PERIODIC_REPORT_STALE_MINUTES=30
|
||||
PERIODIC_REPORT_MAX_ATTEMPTS=3
|
||||
```
|
||||
|
||||
## 多模型分流
|
||||
|
||||
模型管理现在区分两个概念:
|
||||
|
||||
- 可用模型:允许后台任务选择,可同时启用多个;
|
||||
- 默认主模型:只能有一个,正式用户聊天、追问改写和检索重排始终使用它。
|
||||
|
||||
周期报告会选择“周期报告”能力已开启的可用模型,主题摘要和成长档案会选择“摘要沉淀”能力已开启的可用模型。若找不到匹配模型,会自动回退默认主模型,不会因为分流配置缺失直接中断任务。
|
||||
|
||||
部署迁移后,旧版本原来启用的模型会自动成为默认主模型。新增其他模型时建议按以下顺序操作:
|
||||
|
||||
1. 保存模型并执行“测试”;
|
||||
2. 加入可用池;
|
||||
3. 只勾选它实际承担的能力;
|
||||
4. 如需让低成本模型承担报告,应取消默认主模型的“周期报告”能力,避免默认模型优先命中;
|
||||
5. 在数据看板“模型使用与成本”中核对实际模型和成本。
|
||||
|
||||
停用或删除唯一默认主模型会被后端拒绝,必须先启用并设置替代主模型。该限制用于避免生产聊天突然变成无模型可用。
|
||||
|
||||
## 回滚原则
|
||||
|
||||
1. 先停止新版本服务。
|
||||
|
||||
@@ -650,6 +650,7 @@ AI 日志增加:
|
||||
#### 开发进度
|
||||
|
||||
- 2026-07-31:一期已新增模型输入/输出千 Token 单价、币种、适用场景和可用能力字段;AI 请求日志记录模型 ID、估算成本、币种、问题类型、知识命中和路由原因;数据看板展示筛选范围内估算成本。暂未自动切换模型,避免影响正式回答稳定性,后续再基于这些字段做模型分流。
|
||||
- 2026-07-31:二期先完成低风险后台任务分流。模型管理支持“多个可用模型 + 一个默认主模型”;周期报告和主题摘要按能力标签选模型,无匹配时回退默认主模型;正式聊天、追问改写和检索重排仍固定使用默认主模型。后台生成会记录实际模型、分流原因、Token、耗时和估算成本,数据看板增加按调用场景、模型和币种拆分的成本明细。固定信息和正式聊天的自动分流暂不启用,待后台任务运行稳定并核对成本后再做。
|
||||
|
||||
---
|
||||
|
||||
|
||||
Reference in New Issue
Block a user