From fc2af8ba87dda8497fd5dadee4e75f153f5beef6 Mon Sep 17 00:00:00 2001 From: Kfir Amar Date: Sun, 15 Mar 2026 19:33:35 +0200 Subject: [PATCH] fix(api): validate pricing sync and task routing routes --- src/app/api/pricing/sync/route.ts | 27 ++++++++--- src/app/api/settings/task-routing/route.ts | 46 +++++++++++++++---- src/shared/validation/schemas.ts | 52 ++++++++++++++++++++++ tests/unit/t06-schema-hardening.test.mjs | 49 ++++++++++++++++++++ 4 files changed, 161 insertions(+), 13 deletions(-) diff --git a/src/app/api/pricing/sync/route.ts b/src/app/api/pricing/sync/route.ts index b987ce3c3f..b60d78fccf 100644 --- a/src/app/api/pricing/sync/route.ts +++ b/src/app/api/pricing/sync/route.ts @@ -7,14 +7,31 @@ */ import { NextRequest, NextResponse } from "next/server"; +import { pricingSyncRequestSchema } from "@/shared/validation/schemas"; +import { isValidationFailure, validateBody } from "@/shared/validation/helpers"; export async function POST(request: NextRequest) { + let rawBody: unknown; try { - const body = await request.json().catch(() => ({})); - const sources = Array.isArray(body.sources) - ? body.sources.filter((s: unknown): s is string => typeof s === "string") - : undefined; - const dryRun = body.dryRun === true; + rawBody = await request.json(); + } catch { + return NextResponse.json( + { + error: { + message: "Invalid request", + details: [{ field: "body", message: "Invalid JSON body" }], + }, + }, + { status: 400 } + ); + } + + try { + const validation = validateBody(pricingSyncRequestSchema, rawBody); + if (isValidationFailure(validation)) { + return NextResponse.json({ error: validation.error }, { status: 400 }); + } + const { sources, dryRun = false } = validation.data; const { syncPricingFromSources } = await import("@/lib/pricingSync"); const result = await syncPricingFromSources({ sources, dryRun }); diff --git a/src/app/api/settings/task-routing/route.ts b/src/app/api/settings/task-routing/route.ts index c8d346f434..d86fbc2c7b 100644 --- a/src/app/api/settings/task-routing/route.ts +++ b/src/app/api/settings/task-routing/route.ts @@ -6,6 +6,8 @@ import { getDefaultTaskModelMap, } from "@omniroute/open-sse/services/taskAwareRouter.ts"; import { updateSettings } from "@/lib/db/settings"; +import { taskRoutingActionSchema, updateTaskRoutingSchema } from "@/shared/validation/schemas"; +import { isValidationFailure, validateBody } from "@/shared/validation/helpers"; /** * GET /api/settings/task-routing @@ -29,15 +31,29 @@ export async function GET() { * Body: { enabled?: boolean, taskModelMap?: { coding?: "...", ... }, detectionEnabled?: boolean } */ export async function PUT(request: Request) { - let rawBody: Record; + let rawBody: unknown; try { rawBody = await request.json(); } catch { - return NextResponse.json({ error: { message: "Invalid JSON body" } }, { status: 400 }); + return NextResponse.json( + { + error: { + message: "Invalid request", + details: [{ field: "body", message: "Invalid JSON body" }], + }, + }, + { status: 400 } + ); } try { - setTaskRoutingConfig(rawBody as any); + const validation = validateBody(updateTaskRoutingSchema, rawBody); + if (isValidationFailure(validation)) { + return NextResponse.json({ error: validation.error }, { status: 400 }); + } + const config = validation.data; + + setTaskRoutingConfig(config); // Persist to database (excluding stats) const { stats, ...persistable } = getTaskRoutingConfig(); @@ -56,15 +72,29 @@ export async function PUT(request: Request) { * For "detect": pass { action: "detect", body: } to test detection */ export async function POST(request: Request) { - let rawBody: any; + let rawBody: unknown; try { rawBody = await request.json(); } catch { - return NextResponse.json({ error: { message: "Invalid JSON body" } }, { status: 400 }); + return NextResponse.json( + { + error: { + message: "Invalid request", + details: [{ field: "body", message: "Invalid JSON body" }], + }, + }, + { status: 400 } + ); } try { - if (rawBody.action === "reset-stats") { + const validation = validateBody(taskRoutingActionSchema, rawBody); + if (isValidationFailure(validation)) { + return NextResponse.json({ error: validation.error }, { status: 400 }); + } + const actionRequest = validation.data; + + if (actionRequest.action === "reset-stats") { resetTaskRoutingStats(); return NextResponse.json({ success: true, @@ -72,9 +102,9 @@ export async function POST(request: Request) { }); } - if (rawBody.action === "detect") { + if (actionRequest.action === "detect") { const { detectTaskType } = await import("@omniroute/open-sse/services/taskAwareRouter.ts"); - const taskType = detectTaskType(rawBody.body || {}); + const taskType = detectTaskType(actionRequest.body || {}); const config = getTaskRoutingConfig(); return NextResponse.json({ taskType, diff --git a/src/shared/validation/schemas.ts b/src/shared/validation/schemas.ts index c2099350f6..09cb5abd72 100644 --- a/src/shared/validation/schemas.ts +++ b/src/shared/validation/schemas.ts @@ -378,6 +378,58 @@ export const resetStatsActionSchema = z.object({ action: z.literal("reset-stats"), }); +const pricingSyncSourceSchema = z.enum(["litellm"]); + +export const pricingSyncRequestSchema = z + .object({ + sources: z.array(pricingSyncSourceSchema).min(1).optional(), + dryRun: z.boolean().optional(), + }) + .strict(); + +const taskRoutingModelMapSchema = z + .object({ + coding: z.string().max(200).optional(), + creative: z.string().max(200).optional(), + analysis: z.string().max(200).optional(), + vision: z.string().max(200).optional(), + summarization: z.string().max(200).optional(), + background: z.string().max(200).optional(), + chat: z.string().max(200).optional(), + }) + .strict(); + +export const updateTaskRoutingSchema = z + .object({ + enabled: z.boolean().optional(), + taskModelMap: taskRoutingModelMapSchema.optional(), + detectionEnabled: z.boolean().optional(), + }) + .strict() + .superRefine((value, ctx) => { + if ( + value.enabled === undefined && + value.taskModelMap === undefined && + value.detectionEnabled === undefined + ) { + ctx.addIssue({ + code: z.ZodIssueCode.custom, + message: "No valid fields to update", + path: [], + }); + } + }); + +export const taskRoutingActionSchema = z.discriminatedUnion("action", [ + resetStatsActionSchema, + z + .object({ + action: z.literal("detect"), + body: jsonObjectSchema.optional(), + }) + .strict(), +]); + export const updateComboDefaultsSchema = z .object({ comboDefaults: comboRuntimeConfigSchema.optional(), diff --git a/tests/unit/t06-schema-hardening.test.mjs b/tests/unit/t06-schema-hardening.test.mjs index 9c246a3ac3..6140cddaa1 100644 --- a/tests/unit/t06-schema-hardening.test.mjs +++ b/tests/unit/t06-schema-hardening.test.mjs @@ -10,6 +10,9 @@ import { v1EmbeddingsSchema, providerChatCompletionSchema, v1CountTokensSchema, + pricingSyncRequestSchema, + updateTaskRoutingSchema, + taskRoutingActionSchema, } from "../../src/shared/validation/schemas.ts"; test("translatorDetectSchema rejects empty body object", () => { @@ -130,3 +133,49 @@ test("v1CountTokensSchema rejects empty messages", () => { }); assert.equal(validation.success, false); }); + +test("pricingSyncRequestSchema rejects unsupported sources", () => { + const validation = validateBody(pricingSyncRequestSchema, { + sources: ["unknown-source"], + }); + assert.equal(validation.success, false); +}); + +test("pricingSyncRequestSchema accepts dryRun-only requests", () => { + const validation = validateBody(pricingSyncRequestSchema, { + dryRun: true, + }); + assert.equal(validation.success, true); +}); + +test("updateTaskRoutingSchema rejects empty payloads", () => { + const validation = validateBody(updateTaskRoutingSchema, {}); + assert.equal(validation.success, false); +}); + +test("updateTaskRoutingSchema accepts partial task routing updates", () => { + const validation = validateBody(updateTaskRoutingSchema, { + enabled: true, + taskModelMap: { + coding: "codex/gpt-5.1-codex", + }, + }); + assert.equal(validation.success, true); +}); + +test("taskRoutingActionSchema rejects unknown actions", () => { + const validation = validateBody(taskRoutingActionSchema, { + action: "noop", + }); + assert.equal(validation.success, false); +}); + +test("taskRoutingActionSchema accepts detect action with object body", () => { + const validation = validateBody(taskRoutingActionSchema, { + action: "detect", + body: { + messages: [{ role: "user", content: "write code" }], + }, + }); + assert.equal(validation.success, true); +});