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

@@ -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; } } 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 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 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(); } 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 ?? ""; quickModelForm.modelName = target.models[0]?.value ?? "";
} }
async function enableModel(id: number) { async function setModelAvailability(row: ModelItem, enabled: number) {
await api.enableModel(id); await api.setModelAvailability(row.id, enabled);
ElMessage.success("模型已启用"); ElMessage.success(enabled === 1 ? "模型已加入可用池" : "模型已停用");
await loadCurrentMenu(); 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) { async function deleteModel(id: number) {
try { try {
await ElMessageBox.confirm("确认删除该模型?删除后不可恢复。", "删除模型", { 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> Token</span><strong>{{ stats?.totalToken ?? 0 }}</strong></div>
<div class="stat"><span>估算成本</span><strong>{{ formatMoney(stats?.estimatedCost, stats?.costCurrency || 'CNY') }}</strong></div> <div class="stat"><span>估算成本</span><strong>{{ formatMoney(stats?.estimatedCost, stats?.costCurrency || 'CNY') }}</strong></div>
</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> <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> </template>
@@ -1322,7 +1360,7 @@ function formatRecordDateTime(value: string, boundary: "start" | "end") {
<template v-if="activeMenu === 'models'"> <template v-if="activeMenu === 'models'">
<div class="page-head inline"> <div class="page-head inline">
<div><h2>模型管理</h2><p>DeepSeekMiniMax 可快速添加其他供应商按 OpenAI 兼容或 Anthropic Messages 协议添加</p></div> <div><h2>模型管理</h2><p>可同时启用多个模型但只保留一个默认主模型正式聊天始终使用主模型报告和摘要按能力标签分流并自动回退主模型</p></div>
<el-button @click="resetModelForm">清空表单</el-button> <el-button @click="resetModelForm">清空表单</el-button>
</div> </div>
@@ -1444,13 +1482,27 @@ function formatRecordDateTime(value: string, boundary: "start" | "end") {
<el-table-column label="估算单价" width="180"> <el-table-column label="估算单价" width="180">
<template #default="{ row }">{{ row.currency || 'CNY' }} {{ row.inputPricePer1k ?? '-' }}/{{ row.outputPricePer1k ?? '-' }}</template> <template #default="{ row }">{{ row.currency || 'CNY' }} {{ row.inputPricePer1k ?? '-' }}/{{ row.outputPricePer1k ?? '-' }}</template>
</el-table-column> </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="baseUrl" label="Base URL" min-width="220" show-overflow-tooltip />
<el-table-column prop="authType" label="鉴权" width="110" /> <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"> <el-table-column label="操作" width="190" fixed="right" align="center">
<template #default="{ row }"> <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> </template>
</el-table-column> </el-table-column>
</el-table> </el-table>

View File

@@ -207,8 +207,10 @@ export const api = {
request<ModelItem>("/admin/model", { method: "POST", body: JSON.stringify(payload) }), request<ModelItem>("/admin/model", { method: "POST", body: JSON.stringify(payload) }),
updateModel: (id: number, payload: Record<string, unknown>) => updateModel: (id: number, payload: Record<string, unknown>) =>
request<ModelItem>(`/admin/model/${id}`, { method: "PUT", body: JSON.stringify(payload) }), request<ModelItem>(`/admin/model/${id}`, { method: "PUT", body: JSON.stringify(payload) }),
enableModel: (modelId: number) => setModelAvailability: (modelId: number, enabled: number) =>
request<null>("/admin/model/enable", { method: "POST", body: JSON.stringify({ modelId }) }), 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) => deleteModel: (modelId: number) =>
request<null>(`/admin/model/${modelId}`, { method: "DELETE" }), request<null>(`/admin/model/${modelId}`, { method: "DELETE" }),
testModel: (modelId: number) => testModel: (modelId: number) =>

View File

@@ -525,12 +525,15 @@ textarea {
gap: 14px; gap: 14px;
} }
.storage-panel { margin-top: 18px; padding: 18px; border: 1px solid #dfe8e5; border-radius: 8px; background: #fff; } .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 { display: flex; justify-content: space-between; align-items: start; gap: 16px; }
.storage-panel-head h3, .storage-panel-head p { margin: 0 0 6px; } .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 { 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 > div { padding: 14px; background: #f6f8f7; border-radius: 8px; }
.storage-grid span, .storage-grid strong { display: block; } .storage-grid span, .storage-grid strong { display: block; }
.storage-grid strong { margin-top: 8px; font-size: 20px; } .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; } .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 { 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; } .retention-panel div span { display: block; margin-top: 4px; color: #667a73; font-size: 12px; }

View File

@@ -29,6 +29,16 @@ export interface DashboardStats {
totalToken: number; totalToken: number;
estimatedCost?: number | null; estimatedCost?: number | null;
costCurrency?: string | 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 { export interface PromptDetail {
@@ -195,6 +205,7 @@ export interface ModelItem {
remark?: string | null; remark?: string | null;
timeoutSecond: number; timeoutSecond: number;
enabled: number; enabled: number;
isDefault: number;
inputPricePer1k?: number | null; inputPricePer1k?: number | null;
outputPricePer1k?: number | null; outputPricePer1k?: number | null;
currency: string; currency: string;

View File

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

View File

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

View File

@@ -49,6 +49,7 @@ class ModelConfig(Base):
remark: Mapped[str | None] = mapped_column(String(255), nullable=True) remark: Mapped[str | None] = mapped_column(String(255), nullable=True)
timeout_second: Mapped[int] = mapped_column(Integer, default=30, nullable=False) timeout_second: Mapped[int] = mapped_column(Integer, default=30, nullable=False)
enabled: Mapped[int] = mapped_column(default=0, 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) 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) 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) currency: Mapped[str] = mapped_column(String(10), default="CNY", nullable=False)

View File

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

View File

@@ -3,7 +3,7 @@ from __future__ import annotations
from datetime import UTC, datetime from datetime import UTC, datetime
from fastapi import HTTPException, status from fastapi import HTTPException, status
from sqlalchemy import func, select from sqlalchemy import func, select, true
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.core.config import get_settings from app.core.config import get_settings
@@ -133,23 +133,71 @@ class AdminDashboardService:
msg_filter.append(ChatMessage.created_at <= end) msg_filter.append(ChatMessage.created_at <= end)
ai_filter.append(AiRequestLog.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( cost_currency = db.scalar(
select(AiRequestLog.currency) 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()) .order_by(AiRequestLog.id.desc())
.limit(1) .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 { return {
"userCount": db.scalar(select(func.count(User.id)).where(and_(*user_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(and_(*session_filter))) or 0, "sessionCount": db.scalar(select(func.count(ChatSession.id)).where(session_where)) or 0,
"messageCount": db.scalar(select(func.count(ChatMessage.id)).where(and_(*msg_filter))) or 0, "messageCount": db.scalar(select(func.count(ChatMessage.id)).where(message_where)) or 0,
"aiRequestCount": db.scalar(select(func.count(AiRequestLog.id)).where(and_(*ai_filter))) 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, "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, "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(and_(*ai_filter))) 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(and_(*ai_filter))) 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(and_(*ai_filter))) or 0), "estimatedCost": float(db.scalar(select(func.coalesce(func.sum(AiRequestLog.estimated_cost), 0)).where(ai_where)) or 0),
"costCurrency": cost_currency, "costCurrency": cost_currency,
"costBreakdown": cost_breakdown,
} }

View File

@@ -16,9 +16,9 @@ class AiRequestLogService:
def write_success( def write_success(
db: Session, db: Session,
*, *,
session_id: int, session_id: int | None,
message_id: int | None, message_id: int | None,
user_id: int, user_id: int | None,
model_name: str, model_name: str,
prompt: str, prompt: str,
knowledge_ids: str, knowledge_ids: str,
@@ -61,9 +61,9 @@ class AiRequestLogService:
def write_failed( def write_failed(
db: Session, db: Session,
*, *,
session_id: int, session_id: int | None,
message_id: int | None, message_id: int | None,
user_id: int, user_id: int | None,
model_name: str | None, model_name: str | None,
prompt: str | None, prompt: str | None,
knowledge_ids: 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.models.user import User
from app.services.entitlement_service import EntitlementService from app.services.entitlement_service import EntitlementService
from app.services.external_errors import ExternalServiceError 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 from app.services.topic_session_service import TopicSessionService
@@ -62,9 +62,16 @@ class GrowthProfileService:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="当前主题还没有可沉淀的对话内容") raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="当前主题还没有可沉淀的对话内容")
conversation = _messages_text(messages) conversation = _messages_text(messages)
prompt = _topic_summary_prompt(topic, conversation) prompt = _topic_summary_prompt(topic, conversation)
model_name = _enabled_model_name(db) model_name = None
try: 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) parsed = _parse_summary_json(raw)
data = parsed if parsed and "summary" in parsed else _fallback_summary(topic, conversation, raw) data = parsed if parsed and "summary" in parsed else _fallback_summary(topic, conversation, raw)
status_value = "success" status_value = "success"
@@ -96,7 +103,13 @@ class GrowthProfileService:
prompt = _growth_profile_prompt(profile, topic_summary) prompt = _growth_profile_prompt(profile, topic_summary)
try: 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) parsed = _parse_summary_json(raw)
data = parsed if parsed and "profileText" in parsed else _fallback_profile(profile, topic_summary, raw) data = parsed if parsed and "profileText" in parsed else _fallback_profile(profile, topic_summary, raw)
except ExternalServiceError: except ExternalServiceError:
@@ -326,10 +339,5 @@ def _limit(text: str, max_len: int) -> str:
return text[:max_len].strip() 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: def _now() -> datetime:
return datetime.now(UTC).replace(tzinfo=None) return datetime.now(UTC).replace(tzinfo=None)

View File

@@ -12,7 +12,6 @@ from sqlalchemy import or_, select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.core.config import get_settings from app.core.config import get_settings
from app.models.ai_config import ModelConfig
from app.services.chat_context_service import ChatContextService from app.services.chat_context_service import ChatContextService
from app.models.knowledge import ( from app.models.knowledge import (
Knowledge, 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_catalog_cache_service import KnowledgeCatalogCacheService
from app.services.knowledge_service import KnowledgeScope from app.services.knowledge_service import KnowledgeScope
from app.services.model_service import _call_configured_model, _system_config_bool 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 from app.services.rag_service import PromptService, RagResult, RetrievedChunk
SAFETY_RULE_VERSION = "minimum-safety-v1" SAFETY_RULE_VERSION = "minimum-safety-v1"
@@ -264,9 +264,7 @@ class KnowledgeAgentService:
"recentHistory": history_text, "recentHistory": history_text,
} }
fallback = cls._fallback_rewrite(question, recent) fallback = cls._fallback_rewrite(question, recent)
model = db.scalar( model = ModelRoutingService.default_model(db)
select(ModelConfig).where(ModelConfig.enabled == 1).order_by(ModelConfig.id.desc()).limit(1)
)
if model is None or _system_config_bool(db, "mock_model_enabled", get_settings().mock_model_enabled): if model is None or _system_config_bool(db, "mock_model_enabled", get_settings().mock_model_enabled):
return fallback, cls._trace( return fallback, cls._trace(
"rewrite_contextual_question", "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: async def _rerank(cls, db: Session, question: str, candidates: list[Candidate], trace: list[dict], started: float) -> None:
if not candidates: if not candidates:
return 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): 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)) 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 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.core.config import get_settings
from app.models.ai_config import ModelConfig, SystemConfig from app.models.ai_config import ModelConfig, SystemConfig
from app.services.external_errors import ExternalServiceError 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.rag_service import NO_HIT_ANSWER, RagResult
from app.services.secret_service import SecretService from app.services.secret_service import SecretService
@@ -24,6 +25,7 @@ class ModelCompletion:
model_name: str model_name: str
input_token: int input_token: int
output_token: int output_token: int
route_reason: str | None = None
class ModelClientService: class ModelClientService:
@@ -83,6 +85,41 @@ class ModelClientService:
rag_result = RagResult(question=prompt, knowledge_scopes=[], chunks=[], prompt=prompt, allow_general_knowledge=True) 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) 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 @staticmethod
async def summarize_or_raise_async( async def summarize_or_raise_async(
db: Session, db: Session,
@@ -99,12 +136,7 @@ class ModelClientService:
@staticmethod @staticmethod
def _get_enabled_model(db: Session) -> ModelConfig | None: def _get_enabled_model(db: Session) -> ModelConfig | None:
return db.scalar( return ModelRoutingService.default_model(db)
select(ModelConfig)
.where(ModelConfig.enabled == 1)
.order_by(ModelConfig.id.desc())
.limit(1)
)
@staticmethod @staticmethod
def test_model(model: ModelConfig) -> dict[str, Any]: def test_model(model: ModelConfig) -> dict[str, Any]:

View File

@@ -7,7 +7,6 @@ from dataclasses import dataclass
from typing import Any from typing import Any
import httpx import httpx
from sqlalchemy import select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.core.config import get_settings from app.core.config import get_settings
@@ -30,6 +29,7 @@ from app.services.model_service import (
_system_and_turn_messages, _system_and_turn_messages,
_system_config_bool, _system_config_bool,
) )
from app.services.model_routing_service import ModelRoutingService
from app.services.rag_service import RagResult from app.services.rag_service import RagResult
@@ -120,12 +120,7 @@ class ModelStreamService:
def _get_enabled_model(db: Session) -> ModelConfig | None: def _get_enabled_model(db: Session) -> ModelConfig | None:
return db.scalar( return ModelRoutingService.default_model(db)
select(ModelConfig)
.where(ModelConfig.enabled == 1)
.order_by(ModelConfig.id.desc())
.limit(1)
)
def _stream_configured_model(model: ModelConfig, rag_result: RagResult) -> Iterator[str]: 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.core.config import get_settings
from app.models.growth import PeriodicReport, TopicSummary, UserGrowthProfile from app.models.growth import PeriodicReport, TopicSummary, UserGrowthProfile
from app.models.user import User 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"] ReportType = Literal["weekly", "monthly", "stage"]
@@ -177,9 +177,14 @@ class PeriodicReportService:
try: try:
prompt = _report_prompt(user=user, report_type=report_type, period_start=period_start, period_end=period_end, summaries=summaries, profile=profile) 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) completion = TrackedGenerationService.generate(
model = ModelClientService._get_enabled_model(db) db,
report.model_name = model.model_name if model else None 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.status = "success"
report.error_message = None report.error_message = None
except Exception as exc: except Exception as exc:

View File

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

View File

@@ -9,6 +9,7 @@ from sqlalchemy.pool import StaticPool
from app.models import Base from app.models import Base
from app.models.ai_config import ModelConfig from app.models.ai_config import ModelConfig
from app.models.logs import AiRequestLog from app.models.logs import AiRequestLog
from app.services.admin_service import AdminDashboardService
from app.services.ai_request_log_service import AiRequestLogService 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 log.question_type == "knowledge_grounded"
assert "命中知识库" in (log.route_reason or "") assert "命中知识库" in (log.route_reason or "")
def test_dashboard_breaks_cost_down_by_actual_scene_and_model():
with _db() as db:
model = ModelConfig(
id=1,
provider="test",
api_type="openai_compatible",
model_name="report-model",
api_url="https://example.com",
api_key="secret",
input_price_per_1k=Decimal("0.002"),
output_price_per_1k=Decimal("0.006"),
currency="CNY",
timeout_second=30,
)
db.add(model)
db.commit()
AiRequestLogService.write_success(
db,
session_id=None,
message_id=None,
user_id=3,
model_id=1,
model_name="report-model",
prompt="weekly report",
knowledge_ids="",
retrieve_count=0,
input_token=1000,
output_token=500,
cost_ms=120,
question_type="background_report",
route_reason="场景分流:周期报告",
)
db.commit()
stats = AdminDashboardService.stats(db)
assert stats["costBreakdown"] == [
{
"scene": "background_report",
"modelName": "report-model",
"requestCount": 1,
"inputToken": 1000,
"outputToken": 500,
"totalToken": 1500,
"estimatedCost": 0.005,
"currency": "CNY",
}
]

View File

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

View File

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

View File

@@ -163,6 +163,25 @@ PERIODIC_REPORT_STALE_MINUTES=30
PERIODIC_REPORT_MAX_ATTEMPTS=3 PERIODIC_REPORT_MAX_ATTEMPTS=3
``` ```
## 多模型分流
模型管理现在区分两个概念:
- 可用模型:允许后台任务选择,可同时启用多个;
- 默认主模型:只能有一个,正式用户聊天、追问改写和检索重排始终使用它。
周期报告会选择“周期报告”能力已开启的可用模型,主题摘要和成长档案会选择“摘要沉淀”能力已开启的可用模型。若找不到匹配模型,会自动回退默认主模型,不会因为分流配置缺失直接中断任务。
部署迁移后,旧版本原来启用的模型会自动成为默认主模型。新增其他模型时建议按以下顺序操作:
1. 保存模型并执行“测试”;
2. 加入可用池;
3. 只勾选它实际承担的能力;
4. 如需让低成本模型承担报告,应取消默认主模型的“周期报告”能力,避免默认模型优先命中;
5. 在数据看板“模型使用与成本”中核对实际模型和成本。
停用或删除唯一默认主模型会被后端拒绝,必须先启用并设置替代主模型。该限制用于避免生产聊天突然变成无模型可用。
## 回滚原则 ## 回滚原则
1. 先停止新版本服务。 1. 先停止新版本服务。

View File

@@ -650,6 +650,7 @@ AI 日志增加:
#### 开发进度 #### 开发进度
- 2026-07-31一期已新增模型输入/输出千 Token 单价、币种、适用场景和可用能力字段AI 请求日志记录模型 ID、估算成本、币种、问题类型、知识命中和路由原因数据看板展示筛选范围内估算成本。暂未自动切换模型避免影响正式回答稳定性后续再基于这些字段做模型分流。 - 2026-07-31一期已新增模型输入/输出千 Token 单价、币种、适用场景和可用能力字段AI 请求日志记录模型 ID、估算成本、币种、问题类型、知识命中和路由原因数据看板展示筛选范围内估算成本。暂未自动切换模型避免影响正式回答稳定性后续再基于这些字段做模型分流。
- 2026-07-31二期先完成低风险后台任务分流。模型管理支持“多个可用模型 + 一个默认主模型”周期报告和主题摘要按能力标签选模型无匹配时回退默认主模型正式聊天、追问改写和检索重排仍固定使用默认主模型。后台生成会记录实际模型、分流原因、Token、耗时和估算成本数据看板增加按调用场景、模型和币种拆分的成本明细。固定信息和正式聊天的自动分流暂不启用待后台任务运行稳定并核对成本后再做。
--- ---