From 0f73d4e432d18c29c794641a672827d6dc015ab8 Mon Sep 17 00:00:00 2001 From: Xiangzhe Date: Thu, 13 Aug 2026 15:49:38 -0300 Subject: [PATCH] feat(ocr): generic dispatch with per-provider transformation and DI poll loop --- open-sse/config/ocrRegistry.ts | 4 +- open-sse/handlers/ocr.ts | 104 ++++++++++++++++++++---- tests/unit/ocr-handler-dispatch.test.ts | 92 +++++++++++++++++++++ 3 files changed, 180 insertions(+), 20 deletions(-) create mode 100644 tests/unit/ocr-handler-dispatch.test.ts diff --git a/open-sse/config/ocrRegistry.ts b/open-sse/config/ocrRegistry.ts index 27e4dcf994..4bcc141d4c 100644 --- a/open-sse/config/ocrRegistry.ts +++ b/open-sse/config/ocrRegistry.ts @@ -94,8 +94,8 @@ export const AZURE_DI_TRANSFORMATION: OcrTransformation = { analyzeResult?: { content?: string; pages?: unknown[] }; }; const pageCount = r.analyzeResult?.pages?.length ?? 1; - // Azure devolve o markdown do documento inteiro em content; espelhamos no shape - // Mistral com uma "página" agregada por padrão (índice 0), preservando pageCount. + // Azure returns the whole-document markdown in `content`; we mirror it into the + // Mistral shape as a single aggregated "page" (index 0), preserving pageCount. return { pages: [{ index: 0, markdown: r.analyzeResult?.content ?? "" }], model: "prebuilt-read", diff --git a/open-sse/handlers/ocr.ts b/open-sse/handlers/ocr.ts index bf0c553ff0..8e672ab56c 100644 --- a/open-sse/handlers/ocr.ts +++ b/open-sse/handlers/ocr.ts @@ -5,21 +5,43 @@ import { CORS_HEADERS } from "../utils/cors.ts"; * Handles POST /v1/ocr (Mistral OCR API format). */ -import { getOcrProvider, parseOcrModel } from "../config/ocrRegistry.ts"; +import { + getOcrProvider, + getOcrTransformation, + parseOcrModel, + OCR_PROVIDERS, +} from "../config/ocrRegistry.ts"; import { errorResponse } from "../utils/error.ts"; import { attachOmniRouteMetaHeaders } from "@/domain/omnirouteResponseMeta"; import { generateRequestId } from "@/shared/utils/requestId"; +const OCR_POLL_MAX_ATTEMPTS = 30; +const OCR_POLL_INTERVAL_MS = 1000; + +const defaultSleep = (ms: number) => new Promise((resolve) => setTimeout(resolve, ms)); + /** * Handle OCR request * + * Dispatches to the per-provider transformation (see `open-sse/config/ocrRegistry.ts`) + * to build the upstream request, then (for async providers like Azure Document + * Intelligence) polls the returned operation URL until it succeeds or fails, + * before normalizing the response into the Mistral OCR shape. + * * @param {Object} options * @param {Object} options.body - JSON body { model, document } - * @param {Object} options.credentials - Provider credentials { apiKey } + * @param {Object} options.credentials - Provider credentials { apiKey, accessToken, baseUrl } + * @param {Function} [options.fetchImpl] - DI hook for tests; defaults to global fetch + * @param {Function} [options.sleepImpl] - DI hook for tests; defaults to a real setTimeout-based sleep * @returns {Response} */ /** @returns {Promise} */ -export async function handleOcr({ body, credentials }) { +export async function handleOcr({ + body, + credentials, + fetchImpl = fetch, + sleepImpl = defaultSleep, +}) { const startTime = Date.now(); if (!body.document) { return errorResponse(400, "document is required"); @@ -31,7 +53,10 @@ export async function handleOcr({ body, credentials }) { const providerConfig = providerId ? getOcrProvider(providerId) : null; if (!providerConfig) { - return errorResponse(400, `No OCR provider found for model "${model}". Available: mistral`); + return errorResponse( + 400, + `No OCR provider found for model "${model}". Available: ${Object.keys(OCR_PROVIDERS).join(", ")}` + ); } const token = credentials?.apiKey || credentials?.accessToken; @@ -39,18 +64,15 @@ export async function handleOcr({ body, credentials }) { return errorResponse(401, `No credentials for OCR provider: ${providerId}`); } + const baseUrl = credentials?.baseUrl || providerConfig.baseUrl; + if (!baseUrl) { + return errorResponse(400, `No base URL configured for OCR provider: ${providerId}`); + } + try { - const res = await fetch(providerConfig.baseUrl, { - method: "POST", - headers: { - "Content-Type": "application/json", - Authorization: `Bearer ${token}`, - }, - body: JSON.stringify({ - ...body, - model: modelId, - }), - }); + const transformation = getOcrTransformation(providerId); + const { url, init } = transformation.buildRequest({ baseUrl, token, body, modelId }); + const res = await fetchImpl(url, init); if (!res.ok) { const errText = await res.text(); @@ -63,7 +85,17 @@ export async function handleOcr({ body, credentials }) { }); } - const data = await res.json(); + const pollUrl = transformation.pollUrl?.(res) ?? null; + let data: unknown; + if (pollUrl) { + const authHeader = buildAuthHeader(providerConfig.authHeader, token); + data = await pollOcrOperation({ pollUrl, authHeader, fetchImpl, sleepImpl }); + if (data instanceof Response) return data; + } else { + data = await res.json(); + } + + const parsed = transformation.parseResponse(data); const headers = new Headers({ ...CORS_HEADERS, "Content-Type": "application/json" }); attachOmniRouteMetaHeaders(headers, { provider: providerId, @@ -72,8 +104,44 @@ export async function handleOcr({ body, credentials }) { latencyMs: Date.now() - startTime, requestId: generateRequestId(), }); - return new Response(JSON.stringify(data), { status: 200, headers }); + return new Response(JSON.stringify(parsed), { status: 200, headers }); } catch (err) { - return errorResponse(500, `OCR request failed: ${err.message}`); + console.error("[OCR]", err); + return errorResponse(500, "OCR request failed"); } } + +/** + * Build the same auth header used for the initial upstream request, so the + * poll GET (e.g. Azure Document Intelligence's Operation-Location) authenticates + * identically. + */ +function buildAuthHeader(authHeader: string, token: string): Record { + if (authHeader === "bearer") { + return { Authorization: `Bearer ${token}` }; + } + return { [authHeader]: token }; +} + +/** + * Poll an async OCR operation (Azure Document Intelligence) until it succeeds or fails. + * + * @returns {Promise} the parsed JSON body on success, or an error Response + */ +async function pollOcrOperation({ pollUrl, authHeader, fetchImpl, sleepImpl }) { + for (let attempt = 0; attempt < OCR_POLL_MAX_ATTEMPTS; attempt++) { + await sleepImpl(OCR_POLL_INTERVAL_MS); + const pollRes = await fetchImpl(pollUrl, { + method: "GET", + headers: authHeader, + }); + const json = await pollRes.json(); + if (json.status === "succeeded") { + return json; + } + if (json.status === "failed") { + return errorResponse(502, "OCR analysis failed"); + } + } + return errorResponse(504, "OCR analysis timed out"); +} diff --git a/tests/unit/ocr-handler-dispatch.test.ts b/tests/unit/ocr-handler-dispatch.test.ts new file mode 100644 index 0000000000..2ad1e0cec9 --- /dev/null +++ b/tests/unit/ocr-handler-dispatch.test.ts @@ -0,0 +1,92 @@ +import { test } from "node:test"; +import assert from "node:assert/strict"; +import { handleOcr } from "../../open-sse/handlers/ocr.ts"; + +function fetchStub( + script: Array<{ status: number; headers?: Record; json?: unknown }> +) { + const calls: Array<{ url: string; init: RequestInit }> = []; + const impl = async (url: string, init: RequestInit) => { + calls.push({ url, init }); + const step = script.shift()!; + return new Response(step.json !== undefined ? JSON.stringify(step.json) : null, { + status: step.status, + headers: { "Content-Type": "application/json", ...(step.headers ?? {}) }, + }); + }; + return { impl, calls }; +} + +const noSleep = async () => {}; + +test("mistral path posts once and returns the upstream body", async () => { + const { impl, calls } = fetchStub([ + { status: 200, json: { pages: [{ index: 0, markdown: "ok" }], model: "mistral-ocr-latest" } }, + ]); + const res = await handleOcr({ + body: { + model: "mistral/mistral-ocr-latest", + document: { type: "image_url", image_url: "https://x/y.png" }, + }, + credentials: { apiKey: "sk" }, + fetchImpl: impl, + sleepImpl: noSleep, + }); + assert.equal(res.status, 200); + assert.equal(calls.length, 1); + const data = await res.json(); + assert.equal(data.pages[0].markdown, "ok"); +}); + +test("azure DI path polls Operation-Location until succeeded", async () => { + const { impl, calls } = fetchStub([ + { status: 202, headers: { "Operation-Location": "https://poll/op/1" } }, + { status: 200, json: { status: "running" } }, + { status: 200, json: { status: "succeeded", analyzeResult: { content: "# md", pages: [{}] } } }, + ]); + const res = await handleOcr({ + body: { + model: "azure-document-intelligence/prebuilt-read", + document: { type: "document_url", document_url: "https://x/d.pdf" }, + }, + credentials: { apiKey: "azkey", baseUrl: "https://r.cognitiveservices.azure.com" }, + fetchImpl: impl, + sleepImpl: noSleep, + }); + assert.equal(res.status, 200); + assert.ok(calls.length >= 3); + const data = await res.json(); + assert.equal(data.pages[0].markdown, "# md"); +}); + +test("unknown model lists available providers dynamically and errors do not leak internals", async () => { + const res = await handleOcr({ + body: { model: "nope/none", document: { type: "image_url", image_url: "https://x" } }, + credentials: { apiKey: "k" }, + fetchImpl: async () => new Response("{}", { status: 200 }), + sleepImpl: noSleep, + }); + assert.equal(res.status, 400); + const body = await res.json(); + assert.ok(body.error.message.includes("azure-document-intelligence")); + assert.ok(!body.error.message.includes("at /")); +}); + +test("azure DI poll returns failed status maps to 502", async () => { + const { impl } = fetchStub([ + { status: 202, headers: { "Operation-Location": "https://poll/op/1" } }, + { status: 200, json: { status: "failed" } }, + ]); + const res = await handleOcr({ + body: { + model: "azure-document-intelligence/prebuilt-read", + document: { type: "document_url", document_url: "https://x/d.pdf" }, + }, + credentials: { apiKey: "azkey", baseUrl: "https://r.cognitiveservices.azure.com" }, + fetchImpl: impl, + sleepImpl: noSleep, + }); + assert.equal(res.status, 502); + const body = await res.json(); + assert.ok(!body.error.message.includes("at /")); +});