From 28e04857fa3e5382da040eeaee2f50056771586a Mon Sep 17 00:00:00 2001 From: Jesse_Chen Date: Thu, 6 Aug 2026 19:49:26 +0800 Subject: [PATCH] feat(models): secure catalog and pin model versions --- .../src/app/api/admin/model-releases/route.ts | 7 + frontend/src/app/api/admin/models/route.ts | 205 ++++++++ frontend/src/app/api/consult/route.ts | 179 +++---- frontend/src/app/api/models/route.ts | 8 +- .../src/app/api/rectification/agent/route.ts | 116 +++-- .../src/lib/admin/model-mutation-handler.ts | 238 +++++++++ .../src/lib/consultation-model-selection.ts | 4 +- frontend/src/lib/model-catalog.ts | 267 ++++++++++ frontend/src/lib/model-provider-policy.ts | 55 ++ frontend/src/lib/public-models.ts | 4 +- frontend/src/lib/session-model-persistence.ts | 1 + frontend/src/mastra/model.ts | 11 +- .../20260806040000_model_configuration.sql | 481 ++++++++++++++++++ 13 files changed, 1426 insertions(+), 150 deletions(-) create mode 100644 frontend/src/app/api/admin/model-releases/route.ts create mode 100644 frontend/src/app/api/admin/models/route.ts create mode 100644 frontend/src/lib/admin/model-mutation-handler.ts create mode 100644 frontend/src/lib/model-catalog.ts create mode 100644 frontend/src/lib/model-provider-policy.ts create mode 100644 frontend/supabase/migrations/20260806040000_model_configuration.sql diff --git a/frontend/src/app/api/admin/model-releases/route.ts b/frontend/src/app/api/admin/model-releases/route.ts new file mode 100644 index 00000000..edadf017 --- /dev/null +++ b/frontend/src/app/api/admin/model-releases/route.ts @@ -0,0 +1,7 @@ +import { NextResponse } from "next/server"; +import { requirePermission } from "@/lib/admin/auth"; +import { pageOffset, queryAdminRows } from "@/lib/admin/database"; +import { adminErrorResponse, invalidQueryResponse, parseListQuery } from "@/lib/admin/http"; +export const runtime="nodejs"; +type Row={id:string;model_id:string;from_version:number|null;to_version:number;action:string;actor_email:string|null;reason:string;request_id:string;created_at:Date;total_count:string}; +export async function GET(request:Request){try{await requirePermission("models.read");const p=parseListQuery(request);if(!p.success)return invalidQueryResponse(p.error.flatten());const q=p.data.q?`%${p.data.q}%`:null;const rows=await queryAdminRows(`select e.id,c.model_id,f.version from_version,t.version to_version,e.action,u.email actor_email,e.reason,e.request_id,e.created_at,count(*) over()::text total_count from public.model_publish_events e join public.model_configs c on c.id=e.config_id left join public.model_config_versions f on f.id=e.from_version_id join public.model_config_versions t on t.id=e.to_version_id left join identity.users u on u.id=e.actor_user_id where ($1::text is null or c.model_id ilike $1 or u.email ilike $1 or e.request_id ilike $1) and ($2::text is null or e.action=$2) order by e.created_at desc limit $3 offset $4`,[q,p.data.status??null,p.data.pageSize,pageOffset(p.data.page,p.data.pageSize)]);return NextResponse.json({data:rows.map(r=>({id:r.id,modelId:r.model_id,fromVersion:r.from_version,toVersion:r.to_version,action:r.action,actorEmail:r.actor_email,reason:r.reason,requestId:r.request_id,createdAt:r.created_at.toISOString()})),total:Number(rows[0]?.total_count??0)});}catch(e){return adminErrorResponse(e)}} diff --git a/frontend/src/app/api/admin/models/route.ts b/frontend/src/app/api/admin/models/route.ts new file mode 100644 index 00000000..b84bb924 --- /dev/null +++ b/frontend/src/app/api/admin/models/route.ts @@ -0,0 +1,205 @@ +import { NextResponse } from "next/server"; +import { z } from "zod"; + +import { requirePermission } from "@/lib/admin/auth"; +import { pageOffset, queryAdminRows } from "@/lib/admin/database"; +import { + adminErrorResponse, + invalidQueryResponse, + parseListQuery, + requestId, + requireAdminMutation, + requireHighRiskAdminMutation, +} from "@/lib/admin/http"; +import { assertAllowedModelProviderUrl, probeAllowedModelProvider } from "@/lib/epay/gateway-policy"; +import { handleAdminModelMutation, type AdminModelMutation } from "@/lib/admin/model-mutation-handler"; +import { + invalidateLanguageModelCatalog, + modelSettingsContainSecrets, + sanitizeModelSettings, +} from "@/lib/model-catalog"; +import { modelProviderSecretValue } from "@/lib/model-provider-policy"; + +export const runtime = "nodejs"; + +const secretRefSchema = z.string().regex(/^env:[A-Z][A-Z0-9_]*$/); +const providerSchema = z.object({ + action: z.literal("saveProvider"), + id: z.string().uuid().nullable().optional(), + code: z.string().regex(/^[a-z][a-z0-9_-]{1,63}$/), + name: z.string().trim().min(1).max(80), + providerType: z.enum(["openai", "openai-compatible"]), + baseUrl: z.string().url().startsWith("https://").nullable().optional(), + secretRef: secretRefSchema.optional(), + enabled: z.boolean(), + reason: z.string().trim().min(1).max(500), +}).strict(); +const settingsSchema = z.record(z.string(), z.unknown()).superRefine((value, context) => { + if (modelSettingsContainSecrets(value)) { + context.addIssue({ code: "custom", message: "settings 不允许包含密钥、令牌或连接凭据" }); + } +}); +const draftSchema = z.object({ + action: z.literal("saveDraft"), + modelId: z.string().regex(/^[a-z0-9][a-z0-9._-]{0,63}$/), + versionId: z.string().uuid().nullable().optional(), + providerId: z.string().uuid(), + label: z.string().trim().min(1).max(60), + description: z.string().trim().max(200).default(""), + providerModel: z.string().trim().min(1).max(160), + modelTier: z.enum(["standard", "premium", "internal"]), + creditCost: z.number().int().positive(), + contextWindow: z.number().int().positive().nullable().optional(), + inputCostMicrousdPerMillion: z.number().int().nonnegative(), + outputCostMicrousdPerMillion: z.number().int().nonnegative(), + enabled: z.boolean(), + isDefault: z.boolean(), + fallbackModelId: z.string().regex(/^[a-z0-9][a-z0-9._-]{0,63}$/).nullable().optional(), + settings: settingsSchema.default({}), + reason: z.string().trim().min(1).max(500), +}).strict(); +const actionSchema = z.discriminatedUnion("action", [ + providerSchema, + draftSchema, + z.object({ action: z.literal("test"), versionId: z.string().uuid() }).strict(), + z.object({ action: z.literal("publish"), versionId: z.string().uuid(), reason: z.string().trim().min(1).max(500) }).strict(), + z.object({ action: z.literal("rollback"), configId: z.string().uuid(), targetVersion: z.number().int().positive(), reason: z.string().trim().min(1).max(500) }).strict(), +]).superRefine((value, context) => { + if (value.action === "saveDraft" && value.isDefault && !value.enabled) { + context.addIssue({ code: "custom", path: ["isDefault"], message: "默认模型必须启用" }); + } +}); + +type ProviderRow = { + id: string; + code: string; + name: string; + provider_type: "openai" | "openai-compatible"; + base_url: string | null; + secret_ref: string; + enabled: boolean; + updated_at: Date; +}; +type ModelRow = { + id: string; + config_id: string; + model_id: string; + version: number; + provider_id: string; + provider_code: string; + label: string; + description: string; + provider_model: string; + model_tier: string; + credit_cost: number; + context_window: number | null; + input_cost: string; + output_cost: string; + enabled: boolean; + is_default: boolean; + fallback_model_id: string | null; + status: string; + settings: Record; + created_at: Date; + published_at: Date | null; + total_count: string; +}; +const modelOutput = (row: ModelRow) => ({ + id: row.id, + configId: row.config_id, + modelId: row.model_id, + version: row.version, + providerId: row.provider_id, + providerCode: row.provider_code, + label: row.label, + description: row.description, + providerModel: row.provider_model, + modelTier: row.model_tier, + creditCost: row.credit_cost, + contextWindow: row.context_window, + inputCostMicrousdPerMillion: Number(row.input_cost), + outputCostMicrousdPerMillion: Number(row.output_cost), + enabled: row.enabled, + isDefault: row.is_default, + fallbackModelId: row.fallback_model_id, + status: row.status, + settings: sanitizeModelSettings(row.settings), + createdAt: row.created_at.toISOString(), + publishedAt: row.published_at?.toISOString() ?? null, +}); + +export async function GET(request: Request) { + try { + await requirePermission("models.read"); + const parsed = parseListQuery(request); + if (!parsed.success) return invalidQueryResponse(parsed.error.flatten()); + const q = parsed.data.q ? `%${parsed.data.q}%` : null; + const [models, providers] = await Promise.all([ + queryAdminRows(` + select v.id,v.config_id,c.model_id,v.version,v.provider_id,p.code provider_code,v.label,v.description, + v.provider_model,v.model_tier,v.credit_cost,v.context_window, + v.input_cost_microusd_per_million::text input_cost, + v.output_cost_microusd_per_million::text output_cost,v.enabled,v.is_default, + v.fallback_model_id,v.status,v.settings,v.created_at,v.published_at,count(*) over()::text total_count + from public.model_config_versions v + join public.model_configs c on c.id=v.config_id + join public.model_providers p on p.id=v.provider_id + where ($1::text is null or c.model_id ilike $1 or v.label ilike $1 or p.code ilike $1) + and ($2::text is null or v.status=$2) + order by v.created_at desc limit $3 offset $4 + `, [q, parsed.data.status ?? null, parsed.data.pageSize, pageOffset(parsed.data.page, parsed.data.pageSize)]), + queryAdminRows("select id,code,name,provider_type,base_url,secret_ref,enabled,updated_at from public.model_providers order by code"), + ]); + return NextResponse.json({ + data: models.map(modelOutput), + total: Number(models[0]?.total_count ?? 0), + providers: providers.map((provider) => ({ + id: provider.id, + code: provider.code, + name: provider.name, + providerType: provider.provider_type, + baseUrl: provider.base_url, + secretConfigured: Boolean(modelProviderSecretValue({ + code: provider.code, + providerType: provider.provider_type, + secretRef: provider.secret_ref, + })), + enabled: provider.enabled, + updatedAt: provider.updated_at.toISOString(), + })), + }); + } catch (error) { + return adminErrorResponse(error); + } +} + +export async function POST(request: Request) { + try { + const body = actionSchema.safeParse(await request.json().catch(() => null)); + if (!body.success) return invalidQueryResponse(body.error.flatten()); + const permission = body.data.action === "publish" + ? "models.publish" + : body.data.action === "rollback" + ? "models.rollback" + : body.data.action === "test" + ? "models.test" + : "models.write"; + const highRisk = body.data.action === "saveProvider" || body.data.action === "publish" || body.data.action === "rollback"; + const session = highRisk + ? await requireHighRiskAdminMutation(request, permission) + : await requireAdminMutation(request, permission); + const rid = requestId(request); + return await handleAdminModelMutation( + body.data as AdminModelMutation, + { actorUserId: session.user.id, requestId: rid }, + { + queryRows: (sql, values) => queryAdminRows>(sql, values), + assertAllowedUrl: (value) => assertAllowedModelProviderUrl(value).then(() => undefined), + probeAllowed: (value, apiKey) => probeAllowedModelProvider(value, apiKey), + invalidateCatalog: invalidateLanguageModelCatalog, + }, + ); + } catch (error) { + return adminErrorResponse(error); + } +} diff --git a/frontend/src/app/api/consult/route.ts b/frontend/src/app/api/consult/route.ts index 691de977..29338c5c 100644 --- a/frontend/src/app/api/consult/route.ts +++ b/frontend/src/app/api/consult/route.ts @@ -7,17 +7,20 @@ import { runConsultationWorkflow, toAgentConsultationContext, } from "@/mastra"; -import { - languageModelConfigurationMessage, - resolveLanguageModel, -} from "@/mastra/model"; +import { languageModelConfigurationMessage } from "@/mastra/model"; import { blocksPromptExtraction } from "@/lib/consult-safety"; import { consultationEntrypointSchema, resolveConsultationQuestion, } from "@/lib/consultation-entrypoint"; -import { CreditRpcError, runCreditRpc } from "@/lib/consultation-billing"; +import { + authorizeUsage, + completeUsage, + CreditRpcError, + releaseUsage, +} from "@/lib/consultation-billing"; import { reserveConsultationModel } from "@/lib/consultation-model-selection"; +import { resolveSessionLanguageModel } from "@/lib/model-catalog"; import { createAdminSupabaseClient } from "@/lib/supabase/admin"; import { createServerSupabaseClient } from "@/lib/supabase/server"; import { streamTextResponse } from "@/lib/stream-text-response"; @@ -45,6 +48,7 @@ export const maxDuration = 60; const chatRequestMetadataSchema = z.object({ requestId: z.string().uuid(), + sessionId: z.string().uuid(), modelId: z.string().trim().min(1).max(64), name: z.string().trim().max(80).optional().default(""), history: z @@ -166,38 +170,6 @@ function rangeBoundaryWorkflowContext( }; } -async function recordModelUsage( - accounting: ReturnType, - userId: string, - requestId: string, - modelId: string, - usage: Promise<{ inputTokens?: number; outputTokens?: number }>, -) { - try { - const resolved = await usage; - const { error } = await accounting - .from("credit_transactions") - .update({ - model: modelId, - input_tokens: Math.max(0, Math.trunc(resolved.inputTokens ?? 0)), - output_tokens: Math.max(0, Math.trunc(resolved.outputTokens ?? 0)), - }) - .eq("user_id", userId) - .eq("transaction_type", "reserve") - .eq("request_id", requestId); - - if (error) - console.warn( - `[billing] unable to record model usage request=${requestId} model=${modelId}`, - ); - } catch (error) { - const reason = error instanceof Error ? error.name : "UnknownError"; - console.warn( - `[billing] unable to read model usage request=${requestId} model=${modelId} reason=${reason}`, - ); - } -} - export async function POST(request: Request) { let supabase: Awaited>; let accounting: ReturnType; @@ -232,6 +204,35 @@ export async function POST(request: Request) { ); } + const { data: chatSession, error: chatSessionError } = await supabase + .from("chat_sessions") + .select("id,model_id,model_config_version,session_type") + .eq("id", parsed.data.sessionId) + .eq("user_id", user.id) + .maybeSingle(); + if (chatSessionError) { + return NextResponse.json( + { error: "暂时无法读取咨询会话", message: "请稍后重试。" }, + { status: 503 }, + ); + } + if (!chatSession || chatSession.session_type !== "consultation") { + return NextResponse.json( + { error: "咨询会话不存在", message: "请重新进入咨询。" }, + { status: 404 }, + ); + } + if (!chatSession.model_id || chatSession.model_id !== parsed.data.modelId) { + return NextResponse.json( + { error: "会话模型已经变化", message: "请刷新会话后重新发送,本次不会扣点。" }, + { status: 409 }, + ); + } + const sessionModel = await resolveSessionLanguageModel( + chatSession.model_id, + chatSession.model_config_version, + ); + if (parsed.data.entrypoint === "birth_time_rectification") { return NextResponse.json( { @@ -383,20 +384,25 @@ export async function POST(request: Request) { return data; }, reserve: () => reserveConsultationModel( - parsed.data.modelId, - resolveLanguageModel, - () => handoffExecution?.billingReused + chatSession.model_id, + (modelId) => sessionModel?.id === modelId ? sessionModel : null, + (model) => handoffExecution?.billingReused ? Promise.resolve({ success: true, credits: handoffExecution.credits ?? null, error_code: null, }) - : runCreditRpc( - accounting, - "begin_consultation_credit", + : authorizeUsage(accounting, { userId, requestId, - ), + featureKey: "chat.standard", + requestedModelId: model.id, + creditCost: model.creditCost, + }).then((result) => ({ + success: result.success, + credits: result.credits, + error_code: result.reason, + })), ), }); } catch (error) { @@ -499,12 +505,7 @@ export async function POST(request: Request) { return; } try { - await runCreditRpc( - accounting, - "cancel_consultation_credit", - userId, - requestId, - ); + await releaseUsage(accounting, userId, requestId, "consultation_cancelled"); } catch (error) { const reason = error instanceof Error ? error.name : "UnknownError"; console.error( @@ -513,17 +514,28 @@ export async function POST(request: Request) { } } - async function complete() { + const usageStartedAt = Date.now(); + async function complete(usage: Promise<{ inputTokens?: number; outputTokens?: number }>) { if (handoffExecution?.status === "ready") { await settleHandoff(true); return; } - const result = await runCreditRpc( - accounting, - "complete_consultation_credit", - userId, - requestId, - ); + const resolved = await usage; + const inputTokens = Math.max(0, Math.trunc(resolved.inputTokens ?? 0)); + const outputTokens = Math.max(0, Math.trunc(resolved.outputTokens ?? 0)); + const costMicrousd = Math.round(( + inputTokens * (selectedModel.inputCostMicrousdPerMillion ?? 0) + + outputTokens * (selectedModel.outputCostMicrousdPerMillion ?? 0) + ) / 1_000_000); + const result = await completeUsage(accounting, userId, requestId, { + eventKey: requestId, + actualModelId: selectedModel.id, + modelConfigVersion: selectedModel.configVersion, + inputTokens, + outputTokens, + costMicrousd, + durationMs: Date.now() - usageStartedAt, + }); if (!result.success) throw new CreditRpcError(result.error_code || "completion_rejected"); } @@ -552,18 +564,9 @@ export async function POST(request: Request) { ].filter(Boolean).join("\n"), }, ]); - const completeAndRecordUsage = async () => { - await complete(); - void recordModelUsage( - accounting, - userId, - requestId, - modelSelection.usageModelId, - result.totalUsage, - ); - }; + const completeWithUsage = () => complete(result.totalUsage); const settleInterrupted = (emitted: boolean) => - settle(emitted ? completeAndRecordUsage : cancel); + settle(emitted ? completeWithUsage : cancel); return streamTextResponse(result.textStream, { transformText: createBirthTimeModeOutputGuard(consultationMode, false), mode: "mastra", @@ -576,8 +579,8 @@ export async function POST(request: Request) { "x-jyotish-missing-layers": "birth-minute", "x-jyotish-birth-time-mode": consultationMode, }, - ...(handoff ? { onFirstOutput: () => settle(completeAndRecordUsage) } : {}), - onComplete: () => settle(completeAndRecordUsage), + ...(handoff ? { onFirstOutput: () => settle(completeWithUsage) } : {}), + onComplete: () => settle(completeWithUsage), onError: (_error, emitted) => settleInterrupted(emitted), onCancel: settleInterrupted, }); @@ -626,18 +629,9 @@ export async function POST(request: Request) { ].filter(Boolean).join("\n"), }, ]); - const completeAndRecordUsage = async () => { - await complete(); - void recordModelUsage( - accounting, - userId, - requestId, - modelSelection.usageModelId, - result.totalUsage, - ); - }; + const completeWithUsage = () => complete(result.totalUsage); const settleInterrupted = (emitted: boolean) => - settle(emitted ? completeAndRecordUsage : cancel); + settle(emitted ? completeWithUsage : cancel); return streamTextResponse(result.textStream, { transformText: createBirthTimeModeOutputGuard("unverified_birth_time", false), mode: "mastra", @@ -650,8 +644,8 @@ export async function POST(request: Request) { "x-jyotish-missing-layers": workflowReceipt.missingLayers, "x-jyotish-birth-time-mode": "unverified_birth_time", }, - onFirstOutput: () => settle(completeAndRecordUsage), - onComplete: () => settle(completeAndRecordUsage), + onFirstOutput: () => settle(completeWithUsage), + onComplete: () => settle(completeWithUsage), onError: (_error, emitted) => settleInterrupted(emitted), onCancel: settleInterrupted, }); @@ -685,18 +679,9 @@ export async function POST(request: Request) { ].filter(Boolean).join("\n"), }, ]); - const completeAndRecordUsage = async () => { - await complete(); - void recordModelUsage( - accounting, - userId, - requestId, - modelSelection.usageModelId, - result.totalUsage, - ); - }; + const completeWithUsage = () => complete(result.totalUsage); const settleInterrupted = (emitted: boolean) => - settle(emitted ? completeAndRecordUsage : cancel); + settle(emitted ? completeWithUsage : cancel); return streamTextResponse(result.textStream, { transformText: createBirthTimeModeOutputGuard( consultationMode, @@ -712,8 +697,8 @@ export async function POST(request: Request) { "x-jyotish-missing-layers": workflowReceipt.missingLayers, "x-jyotish-birth-time-mode": consultationMode, }, - ...(handoff ? { onFirstOutput: () => settle(completeAndRecordUsage) } : {}), - onComplete: () => settle(completeAndRecordUsage), + ...(handoff ? { onFirstOutput: () => settle(completeWithUsage) } : {}), + onComplete: () => settle(completeWithUsage), onError: (_error, emitted) => settleInterrupted(emitted), onCancel: settleInterrupted, }); @@ -721,7 +706,7 @@ export async function POST(request: Request) { await cancel(); const reason = error instanceof Error ? error.name : "UnknownError"; console.error( - `[consult] generation failed request=${requestId} model=${modelSelection.usageModelId} reason=${reason}`, + `[consult] generation failed request=${requestId} model=${selectedModel.id} reason=${reason}`, ); return NextResponse.json( { diff --git a/frontend/src/app/api/models/route.ts b/frontend/src/app/api/models/route.ts index 11c80f33..7aeb65af 100644 --- a/frontend/src/app/api/models/route.ts +++ b/frontend/src/app/api/models/route.ts @@ -1,6 +1,6 @@ import { NextResponse } from "next/server"; import { createServerSupabaseClient } from "@/lib/supabase/server"; -import { publicLanguageModelCatalog } from "@/mastra/model"; +import { loadLanguageModelCatalog } from "@/lib/model-catalog"; export const runtime = "nodejs"; @@ -26,7 +26,11 @@ export async function GET() { ); } - const catalog = publicLanguageModelCatalog(); + const resolved = await loadLanguageModelCatalog(); + const catalog = { + models: resolved.publicModels, + defaultModelId: resolved.defaultModelId, + }; if (!catalog.defaultModelId || catalog.models.length === 0) { return NextResponse.json( { error: "模型服务尚未配置", message: "当前没有可用的咨询模型。" }, diff --git a/frontend/src/app/api/rectification/agent/route.ts b/frontend/src/app/api/rectification/agent/route.ts index 387790da..41905dca 100644 --- a/frontend/src/app/api/rectification/agent/route.ts +++ b/frontend/src/app/api/rectification/agent/route.ts @@ -3,9 +3,9 @@ import { z } from "zod"; import { parseAgentReply } from "@/lib/agent-reply"; import type { ChatMessage } from "@/lib/chat-message-view"; import { getAgenticRectificationAgent } from "@/mastra/agentic-rectification"; -import { defaultLanguageModel, resolveLanguageModel } from "@/mastra/model"; import { blocksPromptExtraction } from "@/lib/consult-safety"; -import { runCreditRpc } from "@/lib/consultation-billing"; +import { authorizeUsage, completeUsage, releaseUsage } from "@/lib/consultation-billing"; +import { resolveSessionLanguageModel } from "@/lib/model-catalog"; import { createAdminSupabaseClient } from "@/lib/supabase/admin"; import { createServerSupabaseClient } from "@/lib/supabase/server"; import { @@ -80,29 +80,26 @@ function currentTimeContext(now = new Date()) { return `服务端当前时间(权威):${now.toISOString()};中国标准时间(UTC+8):${chinaTime}。涉及“现在、今天、今年、未来几个月”等相对时间时,以此为准。`; } -async function recordModelUsage( +async function rectificationBillingRequestId( accounting: ReturnType, userId: string, - requestId: string, - modelId: string, - usage: Promise<{ inputTokens?: number; outputTokens?: number }>, + sessionId: string, ) { - try { - const resolved = await usage; - const { error } = await accounting - .from("credit_transactions") - .update({ - model: modelId, - input_tokens: Math.max(0, Math.trunc(resolved.inputTokens ?? 0)), - output_tokens: Math.max(0, Math.trunc(resolved.outputTokens ?? 0)), - }) - .eq("user_id", userId) - .eq("transaction_type", "reserve") - .eq("request_id", requestId); - if (error) console.warn(`[agentic-rectification] unable to record usage request=${requestId}`); - } catch (error) { - console.warn(`[agentic-rectification] usage read failed request=${requestId}`, error instanceof Error ? error.name : "UnknownError"); - } + const billingRequestPrefix = `rectification:${sessionId}`; + const { data, error } = await accounting + .from("usage_reservations") + .select("request_id,status") + .eq("user_id", userId) + .eq("feature_key", "rectification") + .like("request_id", `${billingRequestPrefix}%`); + if (error) throw new Error("RectificationBillingLookupError"); + const reservations = (data ?? []) as Array<{ request_id: string; status: string }>; + const active = reservations.find((item) => item.status === "completed") + ?? reservations.find((item) => item.status === "reserved"); + if (active) return active.request_id; + return reservations.length === 0 + ? billingRequestPrefix + : `${billingRequestPrefix}:retry:${reservations.length}`; } export async function GET(request: Request) { @@ -183,7 +180,7 @@ export async function POST(request: Request) { const { data: chatSession, error: chatSessionError } = await supabase .from("chat_sessions") - .select("id,messages,session_type") + .select("id,messages,session_type,model_id,model_config_version") .eq("id", parsed.data.sessionId) .eq("user_id", userId) .maybeSingle(); @@ -251,8 +248,16 @@ export async function POST(request: Request) { ); } - const selectedModel = (conversation.modelId ? resolveLanguageModel(conversation.modelId) : null) - ?? defaultLanguageModel(); + if (!chatSession.model_id || (conversation.modelId && conversation.modelId !== chatSession.model_id)) { + return NextResponse.json( + { error: "会话模型已经变化", message: "请刷新生时校正会话后重试,本次不会扣除点数。" }, + { status: 409 }, + ); + } + const selectedModel = await resolveSessionLanguageModel( + chatSession.model_id, + chatSession.model_config_version, + ); if (!selectedModel) { return NextResponse.json( { error: "模型暂不可用", message: "请选择其他模型后重新发送,本次不会扣除点数。" }, @@ -261,13 +266,20 @@ export async function POST(request: Request) { } let reserveResult; + let billingRequestId: string; try { - reserveResult = await runCreditRpc( + billingRequestId = await rectificationBillingRequestId( accounting, - "begin_consultation_credit", userId, - requestId, + conversation.sessionId, ); + reserveResult = await authorizeUsage(accounting, { + userId, + requestId: billingRequestId, + featureKey: "rectification", + requestedModelId: selectedModel.id, + creditCost: selectedModel.creditCost, + }); } catch (error) { const reason = error instanceof Error ? error.name : "UnknownError"; console.error(`[agentic-rectification] credit reserve failed request=${requestId} reason=${reason}`); @@ -277,11 +289,11 @@ export async function POST(request: Request) { ); } if (!reserveResult.success) { - const insufficient = reserveResult.error_code === "insufficient_credits"; + const insufficient = reserveResult.reason === "insufficient_credits"; return NextResponse.json( { error: insufficient ? "咨询点数不足" : "暂时无法扣除咨询点数", - message: insufficient ? "请先兑换咨询点数后再继续。" : reserveResult.error_code || "请稍后重试。", + message: insufficient ? "请先兑换咨询点数后再继续。" : reserveResult.reason || "请稍后重试。", }, { status: insufficient ? 402 : 503 }, ); @@ -296,18 +308,37 @@ export async function POST(request: Request) { let emitted = false; let raw = ""; let settled = false; - const settle = async (complete: boolean) => { - if (settled) return; + const usageStartedAt = Date.now(); + const settle = async (complete: boolean, usage?: Promise<{ inputTokens?: number; outputTokens?: number }>) => { + if (settled) return true; settled = true; try { + let settlement; if (complete) { - await runCreditRpc(accounting, "complete_consultation_credit", userId, requestId); + const resolved = await usage; + const inputTokens = Math.max(0, Math.trunc(resolved?.inputTokens ?? 0)); + const outputTokens = Math.max(0, Math.trunc(resolved?.outputTokens ?? 0)); + settlement = await completeUsage(accounting, userId, billingRequestId, { + eventKey: requestId, + actualModelId: selectedModel.id, + modelConfigVersion: selectedModel.configVersion, + inputTokens, + outputTokens, + costMicrousd: Math.round(( + inputTokens * (selectedModel.inputCostMicrousdPerMillion ?? 0) + + outputTokens * (selectedModel.outputCostMicrousdPerMillion ?? 0) + ) / 1_000_000), + durationMs: Date.now() - usageStartedAt, + }); } else { - await runCreditRpc(accounting, "cancel_consultation_credit", userId, requestId); + settlement = await releaseUsage(accounting, userId, billingRequestId, "rectification_cancelled"); } + if (!settlement.success) throw new Error(settlement.error_code ?? "usage_settlement_failed"); + return true; } catch (error) { - const reason = error instanceof Error ? error.name : "UnknownError"; - console.warn(`[agentic-rectification] credit settle failed request=${requestId} complete=${complete} reason=${reason}`); + const reason = error instanceof Error ? error.message : "UnknownError"; + console.warn(`[agentic-rectification] usage settle failed request=${requestId} complete=${complete} reason=${reason}`); + return false; } }; const send = (event: Record) => { @@ -335,13 +366,6 @@ export async function POST(request: Request) { raw += chunk; send({ type: "delta", text: chunk }); } - void recordModelUsage( - accounting, - userId, - requestId, - selectedModel.id, - result.totalUsage, - ); const reply = parseAgentReply(raw, "general"); if (!emitted || !reply.text) { console.warn(`[agentic-rectification] empty response request=${requestId}`); @@ -379,8 +403,12 @@ export async function POST(request: Request) { } catch { console.warn(`[agentic-rectification] unable to read candidate result request=${requestId}`); } + if (!await settle(true, result.totalUsage)) { + send({ type: "error", message: "生时校正回复已生成,但用量结算失败,请稍后重试。" }); + controller.close(); + return; + } send({ type: "done", emitted: true }); - await settle(true); controller.close(); } catch (error) { const reason = error instanceof Error ? error.name : "UnknownError"; diff --git a/frontend/src/lib/admin/model-mutation-handler.ts b/frontend/src/lib/admin/model-mutation-handler.ts new file mode 100644 index 00000000..37144505 --- /dev/null +++ b/frontend/src/lib/admin/model-mutation-handler.ts @@ -0,0 +1,238 @@ +import { + expectedModelProviderSecretRef, + modelConnectionTestSucceeded, + modelProviderModelsUrl, + modelProviderSecretValue, + type ModelProviderType, +} from "../model-provider-policy.ts"; + +type SaveProviderAction = { + action: "saveProvider"; + id?: string | null; + code: string; + name: string; + providerType: ModelProviderType; + baseUrl?: string | null; + secretRef?: string; + enabled: boolean; + reason: string; +}; + +type SaveDraftAction = { + action: "saveDraft"; + modelId: string; + versionId?: string | null; + providerId: string; + label: string; + description: string; + providerModel: string; + modelTier: "standard" | "premium" | "internal"; + creditCost: number; + contextWindow?: number | null; + inputCostMicrousdPerMillion: number; + outputCostMicrousdPerMillion: number; + enabled: boolean; + isDefault: boolean; + fallbackModelId?: string | null; + settings: Record; + reason: string; +}; + +export type AdminModelMutation = SaveProviderAction | SaveDraftAction | { + action: "test"; + versionId: string; +} | { + action: "publish"; + versionId: string; + reason: string; +} | { + action: "rollback"; + configId: string; + targetVersion: number; + reason: string; +}; + +type ProviderRow = { + id: string; + code: string; + provider_type: ModelProviderType; + base_url: string | null; + secret_ref: string; + enabled: boolean; + version_id: string; + version_enabled: boolean; +}; + +type Dependencies = { + queryRows(sql: string, values?: readonly unknown[]): Promise[]>; + assertAllowedUrl(value: string): Promise; + probeAllowed(value: string, apiKey: string): Promise; + invalidateCatalog(): void; + environment?: Readonly>; +}; + +function json(body: unknown, status = 200) { + return Response.json(body, { status }); +} + +function isImmutableProviderError(error: unknown) { + return error instanceof Error && error.message.includes("model_provider_runtime_immutable"); +} + +async function versionProvider( + dependencies: Dependencies, + where: string, + values: readonly unknown[], +) { + const rows = await dependencies.queryRows(` + select p.id,p.code,p.provider_type,p.base_url,p.secret_ref,p.enabled, + v.id version_id,v.enabled version_enabled + from public.model_config_versions v + join public.model_providers p on p.id=v.provider_id + where ${where} + `, values); + return (rows[0] as ProviderRow | undefined) ?? null; +} + +async function runnableProvider( + provider: ProviderRow, + dependencies: Dependencies, +) { + if (!provider.enabled || !provider.version_enabled) return null; + const apiKey = modelProviderSecretValue({ + code: provider.code, + providerType: provider.provider_type, + secretRef: provider.secret_ref, + }, dependencies.environment ?? process.env); + if (!apiKey) return null; + if (provider.provider_type === "openai-compatible") { + await dependencies.assertAllowedUrl(provider.base_url ?? ""); + } + return apiKey; +} + +async function probeAndRecord( + provider: ProviderRow, + actorUserId: string, + requestId: string, + dependencies: Dependencies, +) { + const apiKey = await runnableProvider(provider, dependencies); + if (!apiKey) return { recorded: false, status: 0, success: false }; + let status = 0; + try { + status = await dependencies.probeAllowed(modelProviderModelsUrl({ + providerType: provider.provider_type, + baseUrl: provider.base_url, + }), apiKey); + } catch { + status = 0; + } + await dependencies.queryRows( + "select public.admin_record_model_connection_test($1,$2,$3,$4) id", + [actorUserId, provider.version_id, status, requestId], + ); + return { recorded: true, status, success: modelConnectionTestSucceeded(status) }; +} + +export async function handleAdminModelMutation( + action: AdminModelMutation, + context: Readonly<{ actorUserId: string; requestId: string }>, + dependencies: Dependencies, +) { + if (action.action === "saveProvider") { + const secretRef = expectedModelProviderSecretRef(action.code, action.providerType); + if (action.secretRef && action.secretRef !== secretRef) { + return json({ error: "模型供应商密钥引用不受允许" }, 400); + } + if (action.providerType === "openai-compatible") { + await dependencies.assertAllowedUrl(action.baseUrl ?? ""); + } + if (action.enabled && !modelProviderSecretValue({ + code: action.code, + providerType: action.providerType, + secretRef, + }, dependencies.environment ?? process.env)) { + return json({ error: "模型供应商密钥未配置" }, 409); + } + let rows: readonly Record[]; + try { + rows = await dependencies.queryRows( + "select public.admin_save_model_provider($1,$2,$3,$4,$5,$6,$7,$8,$9,$10) id", + [context.actorUserId, action.id ?? null, action.code, action.name, action.providerType, + action.providerType === "openai" ? null : action.baseUrl ?? null, + secretRef, action.enabled, action.reason, context.requestId], + ); + } catch (error) { + if (isImmutableProviderError(error)) { + return json({ + error: "已发布或已退役版本使用的供应商连接配置不可修改,请新建供应商和模型版本后重新测试并发布", + code: "model_provider_runtime_immutable", + }, 409); + } + throw error; + } + return json({ data: { id: rows[0]!.id, requestId: context.requestId } }); + } + + if (action.action === "saveDraft") { + const rows = await dependencies.queryRows( + "select public.admin_save_model_draft($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16::jsonb,$17,$18) id", + [context.actorUserId, action.modelId, action.versionId ?? null, action.providerId, action.label, + action.description, action.providerModel, action.modelTier, action.creditCost, action.contextWindow ?? null, + action.inputCostMicrousdPerMillion, action.outputCostMicrousdPerMillion, action.enabled, action.isDefault, + action.fallbackModelId ?? null, JSON.stringify(action.settings), action.reason, context.requestId], + ); + return json({ data: { id: rows[0]!.id, requestId: context.requestId } }); + } + + if (action.action === "test") { + const provider = await versionProvider( + dependencies, + "v.id=$1 and v.status in ('draft','published','retired')", + [action.versionId], + ); + if (!provider) return json({ error: "模型版本不存在" }, 404); + const result = await probeAndRecord(provider, context.actorUserId, context.requestId, dependencies); + if (!result.recorded) return json({ error: "模型版本或供应商未启用,或部署密钥不可用" }, 409); + return json({ + data: { + id: provider.version_id, + reachable: result.success, + status: result.status, + secretConfigured: true, + requestId: context.requestId, + }, + }, result.success ? 200 : 409); + } + + if (action.action === "publish") { + const rows = await dependencies.queryRows( + "select public.admin_publish_model($1,$2,$3,$4) id", + [context.actorUserId, action.versionId, action.reason, context.requestId], + ); + dependencies.invalidateCatalog(); + return json({ data: { id: rows[0]!.id, requestId: context.requestId } }); + } + + const provider = await versionProvider( + dependencies, + "v.config_id=$1 and v.version=$2 and v.status='retired'", + [action.configId, action.targetVersion], + ); + if (!provider) return json({ error: "回滚版本不存在" }, 404); + const result = await probeAndRecord(provider, context.actorUserId, context.requestId, dependencies); + if (!result.recorded) return json({ error: "回滚版本或供应商未启用,或部署密钥不可用" }, 409); + if (!result.success) { + return json({ + error: "回滚版本运行时连接测试失败", + data: { status: result.status, requestId: context.requestId }, + }, 409); + } + const rows = await dependencies.queryRows( + "select public.admin_rollback_model($1,$2,$3,$4,$5) id", + [context.actorUserId, action.configId, action.targetVersion, action.reason, context.requestId], + ); + dependencies.invalidateCatalog(); + return json({ data: { id: rows[0]!.id, requestId: context.requestId } }); +} diff --git a/frontend/src/lib/consultation-model-selection.ts b/frontend/src/lib/consultation-model-selection.ts index 168fc7c5..c08980c8 100644 --- a/frontend/src/lib/consultation-model-selection.ts +++ b/frontend/src/lib/consultation-model-selection.ts @@ -4,7 +4,7 @@ export async function reserveConsultationModel< >( modelId: string, resolveModel: (modelId: string) => Model | null, - reserveCredit: () => Promise, + reserveCredit: (model: Model) => Promise, ) { const model = resolveModel(modelId); if (!model) return { status: "unavailable" } as const; @@ -13,6 +13,6 @@ export async function reserveConsultationModel< status: "reserved", model, usageModelId: model.id, - reservation: await reserveCredit(), + reservation: await reserveCredit(model), } as const; } diff --git a/frontend/src/lib/model-catalog.ts b/frontend/src/lib/model-catalog.ts new file mode 100644 index 00000000..c468927d --- /dev/null +++ b/frontend/src/lib/model-catalog.ts @@ -0,0 +1,267 @@ +import "server-only"; + +import type { MastraModelConfig } from "@mastra/core/llm"; + +import { queryAdminRows } from "@/lib/admin/database"; +import { loadRuntimeFeatureFlags } from "@/lib/feature-flags"; +import { assertAllowedModelProviderUrl } from "@/lib/epay/gateway-policy"; +import { modelProviderSecretValue, type ModelProviderType } from "@/lib/model-provider-policy"; +import { + languageModelCatalog as environmentCatalog, + type LanguageModelCatalog, + type ResolvedLanguageModel, +} from "@/mastra/model"; + +const cacheTtlMs = 15_000; +const secretSettingNames = new Set([ + "apikey", + "authorization", + "cookie", + "databaseurl", + "connectionstring", + "password", + "privatekey", + "clientsecret", + "secret", + "token", +]); + +type PublishedModelRow = { + model_id: string; + version: number; + label: string; + description: string; + provider_model: string; + credit_cost: number; + is_default: boolean; + provider_code: string; + provider_type: ModelProviderType; + base_url: string | null; + secret_ref: string; + input_cost: string | number; + output_cost: string | number; +}; + +type Cache = { expiresAt: number; catalog: LanguageModelCatalog }; +type Circuit = { failures: number; openUntil: number }; +const state = globalThis as typeof globalThis & { + jyotishaModelCatalogCache?: Cache; + jyotishaModelCatalogCircuit?: Circuit; +}; + +function isSecretSettingKey(key: string) { + const normalized = key.toLowerCase().replace(/[^a-z0-9]/g, ""); + return secretSettingNames.has(normalized) + || normalized.endsWith("apikey") + || normalized.endsWith("password") + || normalized.endsWith("authorization") + || normalized.endsWith("cookie") + || normalized.endsWith("databaseurl") + || normalized.endsWith("connectionstring") + || normalized.endsWith("privatekey") + || normalized.endsWith("clientsecret") + || normalized.endsWith("accesstoken") + || normalized.endsWith("refreshtoken") + || normalized.endsWith("authtoken") + || normalized.endsWith("bearertoken"); +} + +export function modelSettingsContainSecrets(value: unknown): boolean { + if (Array.isArray(value)) return value.some(modelSettingsContainSecrets); + if (!value || typeof value !== "object") return false; + return Object.entries(value).some(([key, nested]) => isSecretSettingKey(key) || modelSettingsContainSecrets(nested)); +} + +export function sanitizeModelSettings(value: unknown): unknown { + if (Array.isArray(value)) return value.map(sanitizeModelSettings); + if (!value || typeof value !== "object") return value; + return Object.fromEntries(Object.entries(value) + .filter(([key]) => !isSecretSettingKey(key)) + .map(([key, nested]) => [key, sanitizeModelSettings(nested)])); +} + +function publicModel(model: ResolvedLanguageModel) { + return { + id: model.id, + label: model.label, + description: model.description, + creditCost: model.creditCost, + isDefault: model.isDefault, + }; +} + +async function resolveRow(row: PublishedModelRow): Promise { + const apiKey = modelProviderSecretValue({ + code: row.provider_code, + providerType: row.provider_type, + secretRef: row.secret_ref, + }); + if (!apiKey) return null; + + const model: MastraModelConfig = row.provider_type === "openai" + ? { + providerId: "openai", + modelId: row.provider_model.replace(/^openai\//, ""), + apiKey, + } + : { + providerId: row.provider_code, + modelId: row.provider_model, + url: (await assertAllowedModelProviderUrl(row.base_url ?? "")).url.toString(), + apiKey, + }; + + return { + id: row.model_id, + label: row.label, + description: row.description, + creditCost: row.credit_cost, + isDefault: row.is_default, + mode: row.provider_type === "openai" ? "openai" : "compatible", + model, + configVersion: row.version, + inputCostMicrousdPerMillion: Number(row.input_cost), + outputCostMicrousdPerMillion: Number(row.output_cost), + }; +} + +async function readPublishedCatalog(): Promise { + const rows = await queryAdminRows(` + select c.model_id,v.version,v.label,v.description,v.provider_model,v.credit_cost,v.is_default, + p.code provider_code,p.provider_type,p.base_url,p.secret_ref, + v.input_cost_microusd_per_million input_cost,v.output_cost_microusd_per_million output_cost + from public.model_config_versions v + join public.model_configs c on c.id=v.config_id + join public.model_providers p on p.id=v.provider_id + where v.status='published' and v.enabled and p.enabled + order by v.is_default desc,c.model_id + `); + if (rows.length === 0) return null; + + const models: ResolvedLanguageModel[] = []; + for (const row of rows) { + try { + const model = await resolveRow(row); + if (model) models.push(model); + } catch { + // Unsafe or unresolvable providers are excluded at the real runtime boundary. + } + } + const defaults = models.filter((model) => model.isDefault); + if (defaults.length !== 1) return null; + return { + models, + publicModels: models.map(publicModel), + defaultModelId: defaults[0]!.id, + issues: [], + }; +} + +function positiveInteger(value: unknown, fallback: number, maximum: number) { + return typeof value === "number" && Number.isInteger(value) && value > 0 + ? Math.min(value, maximum) + : fallback; +} + +async function runtimeSafeEnvironmentCatalog(): Promise { + const models: ResolvedLanguageModel[] = []; + const issues = [...environmentCatalog.issues]; + for (const model of environmentCatalog.models) { + if (model.mode !== "compatible") { + models.push(model); + continue; + } + const config = model.model; + const url = typeof config === "object" && config !== null && "url" in config && typeof config.url === "string" + ? config.url + : ""; + try { + await assertAllowedModelProviderUrl(url); + models.push(model); + } catch { + issues.push(`runtime_provider_unsafe:${model.id}`); + } + } + const defaultModelId = models.some((model) => model.id === environmentCatalog.defaultModelId) + ? environmentCatalog.defaultModelId + : null; + if (!defaultModelId && !issues.includes("default_model_unavailable")) issues.push("default_model_unavailable"); + return { models, publicModels: models.map(publicModel), defaultModelId, issues }; +} + +function recordCatalogFailure(threshold: number, cooldownMs: number) { + const circuit = state.jyotishaModelCatalogCircuit ?? { failures: 0, openUntil: 0 }; + const failures = circuit.failures + 1; + state.jyotishaModelCatalogCircuit = { + failures, + openUntil: failures >= threshold ? Date.now() + cooldownMs : 0, + }; +} + +export async function loadLanguageModelCatalog(): Promise { + const cached = state.jyotishaModelCatalogCache; + if (cached && cached.expiresAt > Date.now()) return cached.catalog; + let catalog = await runtimeSafeEnvironmentCatalog(); + if (process.env.AUTH_PROVIDER?.trim() !== "self-hosted") { + state.jyotishaModelCatalogCache = { expiresAt: Date.now() + cacheTtlMs, catalog }; + return catalog; + } + + let breakerEnabled = false; + let threshold = 5; + let cooldownMs = 300_000; + try { + const flags = await loadRuntimeFeatureFlags(["models.database_catalog", "models.circuit_breaker"]); + if (flags.get("models.database_catalog")?.enabled) { + const breaker = flags.get("models.circuit_breaker"); + breakerEnabled = Boolean(breaker?.enabled); + threshold = positiveInteger(breaker?.config.failureThreshold, threshold, 100); + cooldownMs = positiveInteger(breaker?.config.cooldownSeconds, 300, 3_600) * 1_000; + const circuit = state.jyotishaModelCatalogCircuit ?? { failures: 0, openUntil: 0 }; + if (!breakerEnabled) state.jyotishaModelCatalogCircuit = { failures: 0, openUntil: 0 }; + if (!breakerEnabled || circuit.openUntil <= Date.now()) { + const published = await readPublishedCatalog(); + if (published) { + catalog = published; + state.jyotishaModelCatalogCircuit = { failures: 0, openUntil: 0 }; + } else if (breakerEnabled) { + recordCatalogFailure(threshold, cooldownMs); + } + } + } + } catch { + if (breakerEnabled) recordCatalogFailure(threshold, cooldownMs); + } + state.jyotishaModelCatalogCache = { expiresAt: Date.now() + cacheTtlMs, catalog }; + return catalog; +} + +async function readVersion(modelId: string, version: number) { + const rows = await queryAdminRows(` + select c.model_id,v.version,v.label,v.description,v.provider_model,v.credit_cost,v.is_default, + p.code provider_code,p.provider_type,p.base_url,p.secret_ref, + v.input_cost_microusd_per_million input_cost,v.output_cost_microusd_per_million output_cost + from public.model_config_versions v + join public.model_configs c on c.id=v.config_id + join public.model_providers p on p.id=v.provider_id + where c.model_id=$1 and v.version=$2 and v.status in ('published','retired') and v.enabled and p.enabled + `, [modelId, version]); + return rows[0] ? resolveRow(rows[0]) : null; +} + +export async function resolveSessionLanguageModel(modelId: string, configVersion: number | null | undefined) { + if (configVersion !== null && configVersion !== undefined) { + if (process.env.AUTH_PROVIDER?.trim() !== "self-hosted") return null; + try { + return await readVersion(modelId, configVersion); + } catch { + return null; + } + } + const catalog = await loadLanguageModelCatalog(); + return catalog.models.find((model) => model.id === modelId) ?? null; +} + +export function invalidateLanguageModelCatalog() { + delete state.jyotishaModelCatalogCache; +} diff --git a/frontend/src/lib/model-provider-policy.ts b/frontend/src/lib/model-provider-policy.ts new file mode 100644 index 00000000..d077d063 --- /dev/null +++ b/frontend/src/lib/model-provider-policy.ts @@ -0,0 +1,55 @@ +export type ModelProviderType = "openai" | "openai-compatible"; + +export type ModelProviderSecret = Readonly<{ + code: string; + providerType: ModelProviderType; + secretRef: string; +}>; + +const fixedModelSecretEnvironmentNames = new Set([ + "OPENAI_API_KEY", + "DEEPSEEK_API_KEY", + "LLM_API_KEY", +]); + +export function isAllowedModelSecretEnvironmentName(name: string) { + return fixedModelSecretEnvironmentNames.has(name) + || /^MODEL_PROVIDER_[A-Z][A-Z0-9_]{0,63}_API_KEY$/.test(name); +} + +export function expectedModelProviderSecretRef(code: string, providerType: ModelProviderType) { + if (providerType === "openai") return "env:OPENAI_API_KEY"; + if (code === "deepseek") return "env:DEEPSEEK_API_KEY"; + if (code === "legacy-compatible") return "env:LLM_API_KEY"; + return `env:MODEL_PROVIDER_${code.toUpperCase().replace(/[^A-Z0-9]+/g, "_")}_API_KEY`; +} + +export function modelProviderSecretValue( + provider: ModelProviderSecret, + environment: Readonly> = process.env, +) { + const expected = expectedModelProviderSecretRef(provider.code, provider.providerType); + if (provider.secretRef !== expected) return ""; + const environmentName = expected.slice(4); + if (!isAllowedModelSecretEnvironmentName(environmentName)) return ""; + return environment[environmentName]?.trim() ?? ""; +} + +export function modelProviderModelsUrl(provider: Readonly<{ + providerType: ModelProviderType; + baseUrl: string | null; +}>) { + if (provider.providerType === "openai") return "https://api.openai.com/v1/models"; + const url = new URL(provider.baseUrl ?? ""); + url.username = ""; + url.password = ""; + url.search = ""; + url.hash = ""; + const basePath = url.pathname.replace(/\/+$/, ""); + url.pathname = basePath.endsWith("/models") ? basePath : `${basePath}/models`; + return url.toString(); +} + +export function modelConnectionTestSucceeded(status: number) { + return status >= 200 && status <= 299; +} diff --git a/frontend/src/lib/public-models.ts b/frontend/src/lib/public-models.ts index 48325bb6..bfd2abbd 100644 --- a/frontend/src/lib/public-models.ts +++ b/frontend/src/lib/public-models.ts @@ -4,7 +4,7 @@ const publicLanguageModelSchema = z.object({ id: z.string().trim().min(1).max(64).regex(/^[a-z0-9][a-z0-9._-]*$/), label: z.string().trim().min(1).max(60), description: z.string().trim().max(100), - creditCost: z.literal(1), + creditCost: z.number().int().positive(), isDefault: z.boolean(), }).strict(); @@ -26,7 +26,7 @@ export type PublicLanguageModel = { readonly id: string; readonly label: string; readonly description: string; - readonly creditCost: 1; + readonly creditCost: number; readonly isDefault: boolean; }; diff --git a/frontend/src/lib/session-model-persistence.ts b/frontend/src/lib/session-model-persistence.ts index cbc3c408..f3ee1276 100644 --- a/frontend/src/lib/session-model-persistence.ts +++ b/frontend/src/lib/session-model-persistence.ts @@ -1,4 +1,5 @@ type SessionModelWrite = { + // model_config_version is server-owned and pinned by the database trigger. readonly values: { readonly model_id: string }; readonly sessionId: string; readonly userId: string; diff --git a/frontend/src/mastra/model.ts b/frontend/src/mastra/model.ts index 4d105276..bc8a1d57 100644 --- a/frontend/src/mastra/model.ts +++ b/frontend/src/mastra/model.ts @@ -1,6 +1,8 @@ import type { MastraModelConfig } from "@mastra/core/llm"; import { z } from "zod"; +import { isAllowedModelSecretEnvironmentName } from "@/lib/model-provider-policy"; + type Environment = Readonly>; type LanguageModelMode = "openai" | "compatible"; @@ -8,13 +10,16 @@ export type PublicLanguageModel = { readonly id: string; readonly label: string; readonly description: string; - readonly creditCost: 1; + readonly creditCost: number; readonly isDefault: boolean; }; export type ResolvedLanguageModel = PublicLanguageModel & { readonly mode: LanguageModelMode; readonly model: MastraModelConfig; + readonly configVersion?: number; + readonly inputCostMicrousdPerMillion?: number; + readonly outputCostMicrousdPerMillion?: number; }; export type LanguageModelCatalog = { @@ -25,14 +30,14 @@ export type LanguageModelCatalog = { }; const modelIdSchema = z.string().trim().min(1).max(64).regex(/^[a-z0-9][a-z0-9._-]*$/); -const apiKeyEnvironmentNameSchema = z.string().regex(/^[A-Z][A-Z0-9_]*$/); +const apiKeyEnvironmentNameSchema = z.string().refine(isAllowedModelSecretEnvironmentName); const sharedCatalogFields = { id: modelIdSchema, label: z.string().trim().min(1).max(60), description: z.string().trim().max(100).default(""), apiKeyEnv: apiKeyEnvironmentNameSchema, model: z.string().trim().min(1).max(120), - creditCost: z.literal(1), + creditCost: z.number().int().positive(), }; const catalogEntrySchema = z.discriminatedUnion("provider", [ z.object({ diff --git a/frontend/supabase/migrations/20260806040000_model_configuration.sql b/frontend/supabase/migrations/20260806040000_model_configuration.sql new file mode 100644 index 00000000..4aecc68b --- /dev/null +++ b/frontend/supabase/migrations/20260806040000_model_configuration.sql @@ -0,0 +1,481 @@ +begin; + +create or replace function public.model_provider_base_url_is_safe(p_value text) +returns boolean language plpgsql immutable set search_path='' +as $$ +declare v_url text:=lower(btrim(coalesce(p_value,''))); v_host text; v_port text; +begin + if char_length(v_url) not between 1 and 2048 or v_url !~ '^https://[^/:?#]+(:[0-9]{1,5})?(/[^?#]*)?$' then return false; end if; + v_host:=substring(v_url from '^https://([^/:?#]+)'); + if v_host is null or position('.' in v_host)=0 or v_host !~ '^[a-z0-9]([a-z0-9.-]*[a-z0-9])?$' + or v_host like '%.%' and (v_host like '%..%' or v_host like '%.-%' or v_host like '%-.%') + or v_host ~ '^[0-9.]+$' or v_host ~ '(^|\.)(localhost|local|internal|lan|home|arpa)$' + then return false; end if; + v_port:=substring(v_url from '^https://[^/:?#]+:([0-9]+)'); + return v_port is null or v_port::integer between 1 and 65535; +exception when others then return false; +end $$; + +create or replace function public.expected_model_provider_secret_ref(p_code text,p_provider_type text) +returns text language plpgsql immutable set search_path='' +as $$ +declare v_code text:=btrim(coalesce(p_code,'')); +begin + if p_provider_type='openai' then return 'env:OPENAI_API_KEY'; end if; + if p_provider_type<>'openai-compatible' or v_code !~ '^[a-z][a-z0-9_-]{1,63}$' then return null; end if; + if v_code='deepseek' then return 'env:DEEPSEEK_API_KEY'; end if; + if v_code='legacy-compatible' then return 'env:LLM_API_KEY'; end if; + return 'env:MODEL_PROVIDER_' || upper(regexp_replace(v_code,'[^a-z0-9]+','_','g')) || '_API_KEY'; +end $$; + +create or replace function public.model_settings_contain_secrets(p_value jsonb) +returns boolean language plpgsql immutable set search_path='' +as $$ +declare v_key text; v_child jsonb; v_normalized text; +begin + if p_value is null then return false; end if; + if jsonb_typeof(p_value)='array' then + for v_child in select value from jsonb_array_elements(p_value) loop + if public.model_settings_contain_secrets(v_child) then return true; end if; + end loop; + elsif jsonb_typeof(p_value)='object' then + for v_key,v_child in select key,value from jsonb_each(p_value) loop + v_normalized:=regexp_replace(lower(v_key),'[^a-z0-9]','','g'); + if v_normalized in ('apikey','authorization','cookie','databaseurl','connectionstring','password','privatekey','clientsecret','secret','token') + or v_normalized ~ '(apikey|password|authorization|cookie|databaseurl|connectionstring|privatekey|clientsecret|accesstoken|refreshtoken|authtoken|bearertoken)$' + or public.model_settings_contain_secrets(v_child) + then return true; end if; + end loop; + end if; + return false; +end $$; + +create table if not exists public.model_providers ( + id uuid primary key default gen_random_uuid(), + code text not null unique check (code ~ '^[a-z][a-z0-9_-]{1,63}$'), + name text not null check (char_length(name) between 1 and 80), + provider_type text not null check (provider_type in ('openai', 'openai-compatible')), + base_url text, + secret_ref text not null check ( + secret_ref=public.expected_model_provider_secret_ref(code,provider_type) + ), + enabled boolean not null default true, + created_by uuid references auth.users(id) on delete set null, + updated_by uuid references auth.users(id) on delete set null, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + check ((provider_type='openai' and base_url is null) + or (provider_type='openai-compatible' and public.model_provider_base_url_is_safe(base_url))) +); + +create table if not exists public.model_configs ( + id uuid primary key default gen_random_uuid(), + model_id text not null unique check (model_id ~ '^[a-z0-9][a-z0-9._-]{0,63}$'), + created_at timestamptz not null default now() +); + +create table if not exists public.model_config_versions ( + id uuid primary key default gen_random_uuid(), + config_id uuid not null references public.model_configs(id) on delete restrict, + version integer not null check (version > 0), + provider_id uuid not null references public.model_providers(id) on delete restrict, + label text not null check (char_length(label) between 1 and 60), + description text not null default '' check (char_length(description) <= 200), + provider_model text not null check (char_length(provider_model) between 1 and 160), + model_tier text not null default 'standard' check (model_tier in ('standard','premium','internal')), + credit_cost integer not null default 1 check (credit_cost > 0), + context_window integer check (context_window is null or context_window > 0), + input_cost_microusd_per_million bigint not null default 0 check (input_cost_microusd_per_million >= 0), + output_cost_microusd_per_million bigint not null default 0 check (output_cost_microusd_per_million >= 0), + enabled boolean not null default true, + is_default boolean not null default false, + fallback_model_id text, + status text not null default 'draft' check (status in ('draft','published','retired')), + settings jsonb not null default '{}'::jsonb check (jsonb_typeof(settings)='object' and not public.model_settings_contain_secrets(settings)), + created_by uuid references auth.users(id) on delete set null, + created_at timestamptz not null default now(), + published_at timestamptz, + retired_at timestamptz, + unique(config_id,version), + check (not is_default or enabled) +); +create unique index if not exists model_versions_one_draft_idx on public.model_config_versions(config_id) where status='draft'; +create unique index if not exists model_versions_one_published_idx on public.model_config_versions(config_id) where status='published'; +create unique index if not exists model_versions_one_default_idx on public.model_config_versions((is_default)) where status='published' and enabled and is_default; + +create or replace function public.prevent_published_model_provider_runtime_mutation() +returns trigger language plpgsql set search_path='' +as $$ +begin + if (new.code,new.provider_type,new.base_url,new.secret_ref,new.enabled) + is distinct from + (old.code,old.provider_type,old.base_url,old.secret_ref,old.enabled) + and exists( + select 1 from public.model_config_versions + where provider_id=old.id and status in ('published','retired') + ) + then + raise exception 'model_provider_runtime_immutable' using + errcode='23514', + hint='Create a new provider and model config version, then test and publish it.'; + end if; + return new; +end $$; + +drop trigger if exists model_providers_prevent_published_runtime_mutation on public.model_providers; +create trigger model_providers_prevent_published_runtime_mutation +before update of code,provider_type,base_url,secret_ref,enabled on public.model_providers +for each row execute function public.prevent_published_model_provider_runtime_mutation(); + +create table if not exists public.model_publish_events ( + id uuid primary key default gen_random_uuid(), + config_id uuid not null references public.model_configs(id) on delete restrict, + from_version_id uuid references public.model_config_versions(id) on delete restrict, + to_version_id uuid not null references public.model_config_versions(id) on delete restrict, + action text not null check (action in ('publish','rollback')), + actor_user_id uuid not null references auth.users(id) on delete restrict, + reason text not null check (char_length(btrim(reason)) between 1 and 500), + request_id text not null, + created_at timestamptz not null default clock_timestamp(), + unique(actor_user_id,request_id,action,config_id) +); + +create or replace function public.model_version_config_hash(p_version_id uuid) +returns text language sql stable security definer set search_path='' +as $$ + select encode(public.digest(convert_to(jsonb_build_object( + 'versionId',v.id, + 'configId',v.config_id, + 'version',v.version, + 'providerId',v.provider_id, + 'label',v.label, + 'description',v.description, + 'providerModel',v.provider_model, + 'modelTier',v.model_tier, + 'creditCost',v.credit_cost, + 'contextWindow',v.context_window, + 'inputCost',v.input_cost_microusd_per_million, + 'outputCost',v.output_cost_microusd_per_million, + 'enabled',v.enabled, + 'isDefault',v.is_default, + 'fallbackModelId',v.fallback_model_id, + 'settings',v.settings, + 'providerCode',p.code, + 'providerType',p.provider_type, + 'providerBaseUrl',p.base_url, + 'providerSecretRef',p.secret_ref, + 'providerEnabled',p.enabled + )::text,'utf8'),'sha256'),'hex') + from public.model_config_versions v + join public.model_providers p on p.id=v.provider_id + where v.id=p_version_id +$$; + +create table if not exists public.model_connection_test_evidence ( + id uuid primary key default gen_random_uuid(), + version_id uuid not null references public.model_config_versions(id) on delete cascade, + provider_id uuid not null references public.model_providers(id) on delete restrict, + config_hash text not null check (config_hash ~ '^[0-9a-f]{64}$'), + http_status integer not null check (http_status between 0 and 599), + actor_user_id uuid not null references auth.users(id) on delete restrict, + request_id text not null check (char_length(btrim(request_id)) between 1 and 200), + tested_at timestamptz not null, + expires_at timestamptz not null, + created_at timestamptz not null default clock_timestamp(), + check (expires_at > tested_at), + unique(actor_user_id,request_id,version_id) +); +create index if not exists model_connection_test_evidence_lookup_idx + on public.model_connection_test_evidence(version_id,tested_at desc); + +create or replace function public.model_connection_test_is_fresh(p_version_id uuid,p_request_id text default null) +returns boolean language sql stable security definer set search_path='' +as $$ + select coalesce(( + select e.http_status between 200 and 299 and e.expires_at>clock_timestamp() + from public.model_connection_test_evidence e + where e.version_id=p_version_id + and e.provider_id=(select v.provider_id from public.model_config_versions v where v.id=p_version_id) + and e.config_hash=public.model_version_config_hash(p_version_id) + and (p_request_id is null or e.request_id=p_request_id) + order by e.tested_at desc,e.id desc + limit 1 + ),false) +$$; + +create or replace function public.admin_record_model_connection_test( + p_actor_user_id uuid,p_version_id uuid,p_http_status integer,p_request_id text +) +returns uuid language plpgsql security definer set search_path='' +as $$ +declare v_provider_id uuid; v_hash text; v_id uuid; v_now timestamptz:=clock_timestamp(); +begin + if not (public.admin_has_permission(p_actor_user_id,'models.test') + or public.admin_has_permission(p_actor_user_id,'models.rollback')) + then raise exception 'admin_permission_denied' using errcode='42501'; end if; + if p_http_status not between 0 and 599 then raise exception 'model_connection_status_invalid' using errcode='22023'; end if; + if char_length(btrim(coalesce(p_request_id,''))) not between 1 and 200 then raise exception 'request_id_invalid' using errcode='22023'; end if; + select provider_id,public.model_version_config_hash(id) into v_provider_id,v_hash + from public.model_config_versions where id=p_version_id; + if v_provider_id is null or v_hash is null then raise exception 'model_version_not_found' using errcode='22023'; end if; + insert into public.model_connection_test_evidence( + version_id,provider_id,config_hash,http_status,actor_user_id,request_id,tested_at,expires_at + ) values( + p_version_id,v_provider_id,v_hash,p_http_status,p_actor_user_id,btrim(p_request_id),v_now,v_now+interval '10 minutes' + ) + on conflict(actor_user_id,request_id,version_id) do update set + provider_id=excluded.provider_id,config_hash=excluded.config_hash,http_status=excluded.http_status, + tested_at=excluded.tested_at,expires_at=excluded.expires_at + returning id into v_id; + return v_id; +end $$; + +create or replace function public.admin_save_model_provider( + p_actor_user_id uuid,p_provider_id uuid,p_code text,p_name text,p_provider_type text,p_base_url text, + p_secret_ref text,p_enabled boolean,p_reason text,p_request_id text +) +returns uuid language plpgsql security definer set search_path='' +as $$ declare v_id uuid; v_expected_secret_ref text; +begin + if not public.admin_has_permission(p_actor_user_id,'models.write') then raise exception 'admin_permission_denied' using errcode='42501'; end if; + if char_length(btrim(coalesce(p_reason,''))) not between 1 and 500 then raise exception 'admin_reason_required' using errcode='22023'; end if; + v_expected_secret_ref:=public.expected_model_provider_secret_ref(p_code,p_provider_type); + if v_expected_secret_ref is null or p_secret_ref is distinct from v_expected_secret_ref then + raise exception 'model_provider_secret_ref_forbidden' using errcode='23514'; + end if; + if (p_provider_type='openai' and nullif(btrim(p_base_url),'') is not null) + or (p_provider_type='openai-compatible' and not public.model_provider_base_url_is_safe(p_base_url)) + then raise exception 'model_provider_url_unsafe' using errcode='23514'; end if; + if p_provider_id is null then + insert into public.model_providers(code,name,provider_type,base_url,secret_ref,enabled,created_by,updated_by) + values(p_code,p_name,p_provider_type,nullif(btrim(p_base_url),''),v_expected_secret_ref,p_enabled,p_actor_user_id,p_actor_user_id) returning id into v_id; + else + update public.model_providers set code=p_code,name=p_name,provider_type=p_provider_type,base_url=nullif(btrim(p_base_url),''), + secret_ref=v_expected_secret_ref,enabled=p_enabled,updated_by=p_actor_user_id,updated_at=clock_timestamp() + where id=p_provider_id returning id into v_id; + if v_id is null then raise exception 'provider_not_found' using errcode='22023'; end if; + end if; + insert into audit.admin_audit_logs(actor_user_id,actor_email,actor_role,action,target_type,target_id,after_value, + request_id,permission_used,reason) + select p_actor_user_id,lower(btrim(u.email)),'admin','models.provider.save','model_provider',v_id, + jsonb_build_object('providerCode',p_code,'credentialConfigured',true,'enabled',p_enabled),p_request_id,'models.write',btrim(p_reason) + from identity.users u where u.id=p_actor_user_id on conflict do nothing; + return v_id; +end $$; + +create or replace function public.admin_save_model_draft( + p_actor_user_id uuid,p_model_id text,p_version_id uuid,p_provider_id uuid,p_label text,p_description text, + p_provider_model text,p_model_tier text,p_credit_cost integer,p_context_window integer, + p_input_cost bigint,p_output_cost bigint,p_enabled boolean,p_is_default boolean,p_fallback_model_id text, + p_settings jsonb,p_reason text,p_request_id text +) +returns uuid language plpgsql security definer set search_path='' +as $$ declare v_config_id uuid; v_version integer; v_id uuid; +begin + if not public.admin_has_permission(p_actor_user_id,'models.write') then raise exception 'admin_permission_denied' using errcode='42501'; end if; + if char_length(btrim(coalesce(p_reason,''))) not between 1 and 500 then raise exception 'admin_reason_required' using errcode='22023'; end if; + if p_is_default and not p_enabled then raise exception 'default_model_disabled' using errcode='23514'; end if; + if public.model_settings_contain_secrets(coalesce(p_settings,'{}'::jsonb)) then raise exception 'model_settings_secret_forbidden' using errcode='23514'; end if; + insert into public.model_configs(model_id) values(p_model_id) on conflict(model_id) do update set model_id=excluded.model_id returning id into v_config_id; + if p_version_id is null then + select coalesce(max(version),0)+1 into v_version from public.model_config_versions where config_id=v_config_id; + insert into public.model_config_versions(config_id,version,provider_id,label,description,provider_model,model_tier, + credit_cost,context_window,input_cost_microusd_per_million,output_cost_microusd_per_million,enabled,is_default, + fallback_model_id,settings,created_by) + values(v_config_id,v_version,p_provider_id,p_label,p_description,p_provider_model,p_model_tier,p_credit_cost,p_context_window, + p_input_cost,p_output_cost,p_enabled,p_is_default,nullif(btrim(p_fallback_model_id),''),coalesce(p_settings,'{}'::jsonb),p_actor_user_id) + returning id into v_id; + else + update public.model_config_versions set provider_id=p_provider_id,label=p_label,description=p_description, + provider_model=p_provider_model,model_tier=p_model_tier,credit_cost=p_credit_cost,context_window=p_context_window, + input_cost_microusd_per_million=p_input_cost,output_cost_microusd_per_million=p_output_cost,enabled=p_enabled, + is_default=p_is_default,fallback_model_id=nullif(btrim(p_fallback_model_id),''),settings=coalesce(p_settings,'{}'::jsonb) + where id=p_version_id and config_id=v_config_id and status='draft' returning id into v_id; + if v_id is null then raise exception 'model_draft_not_found' using errcode='22023'; end if; + end if; + insert into audit.admin_audit_logs(actor_user_id,actor_email,actor_role,action,target_type,target_id,after_value, + request_id,permission_used,reason) + select p_actor_user_id,lower(btrim(u.email)),'admin','models.draft.save','model_config_version',v_id, + jsonb_build_object('modelId',p_model_id,'providerId',p_provider_id,'enabled',p_enabled,'isDefault',p_is_default), + p_request_id,'models.write',btrim(p_reason) from identity.users u where u.id=p_actor_user_id on conflict do nothing; + return v_id; +end $$; + +create or replace function public.admin_publish_model( + p_actor_user_id uuid,p_version_id uuid,p_reason text,p_request_id text +) +returns uuid language plpgsql security definer set search_path='' +as $$ +declare v_draft public.model_config_versions%rowtype; v_model_id text; v_from uuid; v_cycle boolean; +begin + if not public.admin_has_permission(p_actor_user_id,'models.publish') then raise exception 'admin_permission_denied' using errcode='42501'; end if; + if char_length(btrim(coalesce(p_reason,''))) not between 1 and 500 then raise exception 'admin_reason_required' using errcode='22023'; end if; + select * into v_draft from public.model_config_versions where id=p_version_id and status='draft' for update; + if not found then raise exception 'model_draft_not_found' using errcode='22023'; end if; + if v_draft.is_default and not v_draft.enabled then raise exception 'default_model_disabled' using errcode='23514'; end if; + if public.model_settings_contain_secrets(v_draft.settings) then raise exception 'model_settings_secret_forbidden' using errcode='23514'; end if; + perform 1 from public.model_providers p where p.id=v_draft.provider_id and p.enabled + and p.secret_ref=public.expected_model_provider_secret_ref(p.code,p.provider_type) + and (p.provider_type='openai' or public.model_provider_base_url_is_safe(p.base_url)) + for update; + if not found then raise exception 'model_provider_unavailable' using errcode='23514'; end if; + if not public.model_connection_test_is_fresh(p_version_id,null) then + raise exception 'model_connection_test_required' using errcode='23514'; + end if; + select model_id into v_model_id from public.model_configs where id=v_draft.config_id; + if v_draft.fallback_model_id=v_model_id then raise exception 'model_fallback_cycle' using errcode='23514'; end if; + if v_draft.fallback_model_id is not null and not exists( + select 1 from public.model_configs c + join public.model_config_versions v on v.config_id=c.id + join public.model_providers p on p.id=v.provider_id + where c.model_id=v_draft.fallback_model_id and v.status='published' and v.enabled and p.enabled + and p.secret_ref=public.expected_model_provider_secret_ref(p.code,p.provider_type) + and (p.provider_type='openai' or public.model_provider_base_url_is_safe(p.base_url)) + ) then raise exception 'fallback_model_unavailable' using errcode='23514'; end if; + with recursive edges(model_id,fallback_model_id) as ( + select c.model_id,case when v.id=p_version_id then v_draft.fallback_model_id else v.fallback_model_id end + from public.model_configs c join public.model_config_versions v on v.config_id=c.id + where (v.status='published' and v.config_id<>v_draft.config_id) or v.id=p_version_id + ), walk(origin,node,path,cycle) as ( + select model_id,fallback_model_id,array[model_id],false from edges where fallback_model_id is not null + union all + select w.origin,e.fallback_model_id,w.path||e.model_id,e.model_id=any(w.path) + from walk w join edges e on e.model_id=w.node where not w.cycle and e.fallback_model_id is not null + ) select coalesce(bool_or(cycle),false) into v_cycle from walk; + if v_cycle then raise exception 'model_fallback_cycle' using errcode='23514'; end if; + select id into v_from from public.model_config_versions where config_id=v_draft.config_id and status='published' for update; + if v_from is not null and not v_draft.is_default and exists( + select 1 from public.model_config_versions where id=v_from and is_default + ) then + if not v_draft.enabled then raise exception 'default_model_disabled' using errcode='23514'; end if; + update public.model_config_versions set is_default=true where id=p_version_id; + v_draft.is_default:=true; + end if; + if v_draft.is_default then update public.model_config_versions set is_default=false where status='published' and is_default; end if; + update public.model_config_versions set status='retired',is_default=false,retired_at=clock_timestamp() where id=v_from; + update public.model_config_versions set status='published',published_at=clock_timestamp(),retired_at=null where id=p_version_id; + if not exists( + select 1 from public.model_config_versions v join public.model_providers p on p.id=v.provider_id + where v.status='published' and v.enabled and v.is_default and p.enabled + and p.secret_ref=public.expected_model_provider_secret_ref(p.code,p.provider_type) + and (p.provider_type='openai' or public.model_provider_base_url_is_safe(p.base_url)) + ) then raise exception 'default_model_required' using errcode='23514'; end if; + insert into public.model_publish_events(config_id,from_version_id,to_version_id,action,actor_user_id,reason,request_id) + values(v_draft.config_id,v_from,p_version_id,'publish',p_actor_user_id,btrim(p_reason),p_request_id) on conflict do nothing; + insert into audit.admin_audit_logs(actor_user_id,actor_email,actor_role,action,target_type,target_id,after_value, + request_id,permission_used,reason) + select p_actor_user_id,lower(btrim(u.email)),'admin','models.publish','model_config_version',p_version_id, + jsonb_build_object('modelId',v_model_id,'version',v_draft.version),p_request_id,'models.publish',btrim(p_reason) + from identity.users u where u.id=p_actor_user_id on conflict do nothing; + return p_version_id; +end $$; + +create or replace function public.admin_rollback_model( + p_actor_user_id uuid,p_config_id uuid,p_target_version integer,p_reason text,p_request_id text +) +returns uuid language plpgsql security definer set search_path='' +as $$ +declare v_current public.model_config_versions%rowtype; v_target public.model_config_versions%rowtype; +begin + if not public.admin_has_permission(p_actor_user_id,'models.rollback') then raise exception 'admin_permission_denied' using errcode='42501'; end if; + if char_length(btrim(coalesce(p_reason,''))) not between 1 and 500 then raise exception 'admin_reason_required' using errcode='22023'; end if; + select * into v_current from public.model_config_versions where config_id=p_config_id and status='published' for update; + select * into v_target from public.model_config_versions where config_id=p_config_id and version=p_target_version and status='retired' for update; + if v_current.id is null or v_target.id is null then raise exception 'rollback_version_not_found' using errcode='22023'; end if; + if not v_target.enabled then raise exception 'rollback_model_disabled' using errcode='23514'; end if; + if public.model_settings_contain_secrets(v_target.settings) then raise exception 'model_settings_secret_forbidden' using errcode='23514'; end if; + if not exists( + select 1 from public.model_providers p where p.id=v_target.provider_id and p.enabled + and p.secret_ref=public.expected_model_provider_secret_ref(p.code,p.provider_type) + and (p.provider_type='openai' or public.model_provider_base_url_is_safe(p.base_url)) + ) then raise exception 'model_provider_unavailable' using errcode='23514'; end if; + if not public.model_connection_test_is_fresh(v_target.id,p_request_id) then + raise exception 'model_connection_test_required' using errcode='23514'; + end if; + v_target.is_default:=v_current.is_default; + if v_target.is_default then update public.model_config_versions set is_default=false where status='published' and is_default; end if; + update public.model_config_versions set status='retired',is_default=false,retired_at=clock_timestamp() where id=v_current.id; + update public.model_config_versions set status='published',is_default=v_target.is_default,published_at=clock_timestamp(),retired_at=null where id=v_target.id; + if not exists( + select 1 from public.model_config_versions v join public.model_providers p on p.id=v.provider_id + where v.status='published' and v.enabled and v.is_default and p.enabled + and p.secret_ref=public.expected_model_provider_secret_ref(p.code,p.provider_type) + and (p.provider_type='openai' or public.model_provider_base_url_is_safe(p.base_url)) + ) then raise exception 'default_model_required' using errcode='23514'; end if; + insert into public.model_publish_events(config_id,from_version_id,to_version_id,action,actor_user_id,reason,request_id) + values(p_config_id,v_current.id,v_target.id,'rollback',p_actor_user_id,btrim(p_reason),p_request_id) on conflict do nothing; + insert into audit.admin_audit_logs(actor_user_id,actor_email,actor_role,action,target_type,target_id,after_value, + request_id,permission_used,reason) + select p_actor_user_id,lower(btrim(u.email)),'admin','models.rollback','model_config_version',v_target.id, + jsonb_build_object('version',p_target_version),p_request_id,'models.rollback',btrim(p_reason) + from identity.users u where u.id=p_actor_user_id on conflict do nothing; + return v_target.id; +end $$; + +alter table public.chat_sessions add column if not exists model_config_version integer; + +update public.chat_sessions s +set model_config_version=v.version +from public.model_configs c +join public.model_config_versions v on v.config_id=c.id and v.status='published' +where s.model_id=c.model_id and s.model_config_version is null; + +create or replace function public.pin_chat_session_model_config_version() +returns trigger language plpgsql security definer set search_path='' +as $$ +begin + if tg_op='UPDATE' and new.model_id is not distinct from old.model_id and old.model_config_version is not null then + new.model_config_version:=old.model_config_version; + return new; + end if; + if new.model_id is null then + new.model_config_version:=null; + return new; + end if; + select v.version into new.model_config_version + from public.model_configs c + join public.model_config_versions v on v.config_id=c.id + where c.model_id=new.model_id and v.status='published' and v.enabled + limit 1; + if not found then new.model_config_version:=null; end if; + return new; +end $$; + +drop trigger if exists chat_sessions_pin_model_config_version on public.chat_sessions; +create trigger chat_sessions_pin_model_config_version +before insert or update of model_id,model_config_version on public.chat_sessions +for each row execute function public.pin_chat_session_model_config_version(); + +alter table public.model_providers enable row level security; +alter table public.model_configs enable row level security; +alter table public.model_config_versions enable row level security; +alter table public.model_publish_events enable row level security; +alter table public.model_connection_test_evidence enable row level security; +revoke all on table public.model_providers,public.model_configs,public.model_config_versions,public.model_publish_events,public.model_connection_test_evidence from public,anon,authenticated; +grant select on table public.model_providers,public.model_configs,public.model_config_versions,public.model_publish_events,public.model_connection_test_evidence to service_role; +revoke all on function public.model_provider_base_url_is_safe(text),public.expected_model_provider_secret_ref(text,text), + public.model_settings_contain_secrets(jsonb),public.prevent_published_model_provider_runtime_mutation(), + public.model_version_config_hash(uuid),public.model_connection_test_is_fresh(uuid,text), + public.admin_record_model_connection_test(uuid,uuid,integer,text), + public.admin_save_model_provider(uuid,uuid,text,text,text,text,text,boolean,text,text), + public.admin_save_model_draft(uuid,text,uuid,uuid,text,text,text,text,integer,integer,bigint,bigint,boolean,boolean,text,jsonb,text,text), + public.admin_publish_model(uuid,uuid,text,text),public.admin_rollback_model(uuid,uuid,integer,text,text), + public.pin_chat_session_model_config_version() + from public,anon,authenticated; +grant execute on function public.model_provider_base_url_is_safe(text),public.expected_model_provider_secret_ref(text,text), + public.model_settings_contain_secrets(jsonb),public.model_version_config_hash(uuid),public.model_connection_test_is_fresh(uuid,text), + public.admin_record_model_connection_test(uuid,uuid,integer,text), + public.admin_save_model_provider(uuid,uuid,text,text,text,text,text,boolean,text,text), + public.admin_save_model_draft(uuid,text,uuid,uuid,text,text,text,text,integer,integer,bigint,bigint,boolean,boolean,text,jsonb,text,text), + public.admin_publish_model(uuid,uuid,text,text),public.admin_rollback_model(uuid,uuid,integer,text,text) to service_role; +do $$ begin if exists(select 1 from pg_roles where rolname='admin_runtime') then + grant select on table public.model_providers,public.model_configs,public.model_config_versions,public.model_publish_events,public.model_connection_test_evidence to admin_runtime; + grant execute on function public.model_provider_base_url_is_safe(text),public.expected_model_provider_secret_ref(text,text), + public.model_settings_contain_secrets(jsonb),public.model_version_config_hash(uuid),public.model_connection_test_is_fresh(uuid,text), + public.admin_record_model_connection_test(uuid,uuid,integer,text), + public.admin_save_model_provider(uuid,uuid,text,text,text,text,text,boolean,text,text), + public.admin_save_model_draft(uuid,text,uuid,uuid,text,text,text,text,integer,integer,bigint,bigint,boolean,boolean,text,jsonb,text,text), + public.admin_publish_model(uuid,uuid,text,text),public.admin_rollback_model(uuid,uuid,integer,text,text) to admin_runtime; +end if; end $$; + +commit;