From ec400a06b7603517eb5225674cf02f88a9014d59 Mon Sep 17 00:00:00 2001 From: Peter Steinberger Date: Sun, 6 Sep 2026 20:12:14 -0700 Subject: [PATCH] refactor(embeddings): share compatible batch and request orchestration (#140660) --- extensions/openai/embedding-batch.ts | 144 +++++------------- extensions/voyage/embedding-batch.test.ts | 124 ++++++++++----- extensions/voyage/embedding-batch.ts | 134 +++++----------- extensions/voyage/embedding-provider.ts | 81 ++-------- .../memory-host-sdk/src/engine-embeddings.ts | 1 + .../memory-host-sdk/src/host/batch-runner.ts | 82 ++++++++++ .../memory-core-host-engine-embeddings.ts | 1 + 7 files changed, 253 insertions(+), 314 deletions(-) diff --git a/extensions/openai/embedding-batch.ts b/extensions/openai/embedding-batch.ts index 2ea4e35d4f6a..cbc2f6db3daf 100644 --- a/extensions/openai/embedding-batch.ts +++ b/extensions/openai/embedding-batch.ts @@ -1,7 +1,6 @@ import { coerceErrorMessage as formatOpenAiBatchError } from "openclaw/plugin-sdk/error-runtime"; // Openai plugin module implements embedding batch behavior. import { - applyEmbeddingBatchOutputLine, buildBatchHeaders, buildEmbeddingBatchGroupOptions, EMBEDDING_BATCH_ENDPOINT, @@ -9,11 +8,8 @@ import { formatBatchErrorDetail, formatUnavailableBatchError, postJsonWithRetry, - readEmbeddingBatchJsonl, resolveEmbeddingEndpointUrl, - resolveCompletedBatchResult, - runEmbeddingBatchGroups, - throwIfBatchCompletionError, + runEmbeddingBatches, type EmbeddingBatchExecutionParams, type EmbeddingBatchStatus, type ProviderBatchOutputLine, @@ -114,25 +110,6 @@ async function fetchOpenAiFileContent(params: { }); } -async function readOpenAiBatchOutputFile(params: { - openAi: OpenAiEmbeddingClient; - fileId: string; - maxLines: number; - onLine: (line: OpenAiBatchOutputLine) => boolean; -}): Promise { - return await fetchOpenAiBatchResource({ - openAi: params.openAi, - path: `/files/${params.fileId}/content`, - label: "openai.batch-file-content", - parse: async (res) => - await readEmbeddingBatchJsonl(res, { - label: "openai.batch-file-content", - maxRecords: params.maxLines, - onRecord: params.onLine, - }), - }); -} - async function fetchOpenAiBatchResource(params: { openAi: OpenAiEmbeddingClient; path: string; @@ -238,7 +215,8 @@ export async function runOpenAiEmbeddingBatches( maxJsonlBytes?: number; } & EmbeddingBatchExecutionParams, ): Promise> { - return await runEmbeddingBatchGroups({ + return await runEmbeddingBatches({ + provider: "openai", ...buildEmbeddingBatchGroupOptions(params, { maxRequests: OPENAI_BATCH_MAX_REQUESTS, maxJsonlBytes: params.maxJsonlBytes ?? OPENAI_BATCH_MAX_JSONL_BYTES, @@ -253,94 +231,46 @@ export async function runOpenAiEmbeddingBatches( error: formatOpenAiBatchDiagnostic(error), }); }, - runGroup: async ({ group, groupIndex, groups, byCustomId, pollIntervalMs, timeoutMs }) => { - const batchInfo = await submitOpenAiBatch({ + submit: (group) => + submitOpenAiBatch({ openAi: params.openAi, requests: group, agentId: params.agentId }), + readError: (errorFileId) => readOpenAiBatchError({ openAi: params.openAi, errorFileId }), + readOutput: (fileId, parse) => + fetchOpenAiBatchResource({ openAi: params.openAi, - requests: group, - agentId: params.agentId, - }); - if (!batchInfo.id) { - throw new Error("openai batch create failed: missing batch id"); - } + path: `/files/${fileId}/content`, + label: "openai.batch-file-content", + parse, + }), + waitForBatch: async (batchInfo, pollIntervalMs, timeoutMs) => { const batchId = batchInfo.id; - - params.debug?.("memory embeddings: openai batch created", { - batchId: batchInfo.id, - status: batchInfo.status, - group: groupIndex + 1, - groups, - requests: group.length, + const openAi = params.openAi; + const wait = params.wait; + const debug = params.debug; + const deadline = createProviderOperationDeadline({ + label: `openai batch ${batchId}`, + timeoutMs, }); - - await throwIfBatchCompletionError({ + return await waitForEmbeddingBatch({ provider: "openai", - status: batchInfo, - readError: async (errorFileId) => - await readOpenAiBatchError({ openAi: params.openAi, errorFileId }), - }); - - const completed = await resolveCompletedBatchResult({ - provider: "openai", - status: batchInfo, - wait: params.wait, - waitForBatch: async () => { - const openAi = params.openAi; - const wait = params.wait; - const debug = params.debug; - const deadline = createProviderOperationDeadline({ - label: `openai batch ${batchId}`, - timeoutMs, - }); - return await waitForEmbeddingBatch({ - provider: "openai", - batchId, - wait, - pollIntervalMs, - timeoutMs, - debug, - initial: batchInfo, - fetchStatus: (signal) => fetchOpenAiBatchStatus({ openAi, batchId, signal }), - resolveTimeoutMs: () => - resolveProviderOperationTimeoutMs({ deadline, defaultTimeoutMs: timeoutMs }), - waitForPoll: (delayMs) => - waitProviderOperationPollInterval({ deadline, pollIntervalMs: delayMs }), - readError: async (errorFileId) => await readOpenAiBatchError({ openAi, errorFileId }), - backoff: { - maxDelayMs: OPENAI_BATCH_MAX_POLL_BACKOFF_MS, - shouldRetry: isRetryableOpenAiBatchPollError, - formatError: formatOpenAiBatchDiagnostic, - formatProgress: formatOpenAiBatchProgress, - }, - }); + batchId, + wait, + pollIntervalMs, + timeoutMs, + debug, + initial: batchInfo, + fetchStatus: (signal) => fetchOpenAiBatchStatus({ openAi, batchId, signal }), + resolveTimeoutMs: () => + resolveProviderOperationTimeoutMs({ deadline, defaultTimeoutMs: timeoutMs }), + waitForPoll: (delayMs) => + waitProviderOperationPollInterval({ deadline, pollIntervalMs: delayMs }), + readError: async (errorFileId) => await readOpenAiBatchError({ openAi, errorFileId }), + backoff: { + maxDelayMs: OPENAI_BATCH_MAX_POLL_BACKOFF_MS, + shouldRetry: isRetryableOpenAiBatchPollError, + formatError: formatOpenAiBatchDiagnostic, + formatProgress: formatOpenAiBatchProgress, }, }); - - const errors: string[] = []; - const remaining = new Set(group.map((request) => request.custom_id)); - - await readOpenAiBatchOutputFile({ - openAi: params.openAi, - fileId: completed.outputFileId, - maxLines: group.length, - onLine: (line) => { - // Only the first response for a submitted id may mutate results. - if (line.custom_id && remaining.has(line.custom_id)) { - applyEmbeddingBatchOutputLine({ line, remaining, errors, byCustomId }); - } - return errors.length === 0 && remaining.size > 0; - }, - }); - - if (errors.length > 0) { - throw new Error( - `openai batch ${batchInfo.id} failed: ${formatBatchErrorDetail(errors[0]) ?? "unknown error"}`, - ); - } - if (remaining.size > 0) { - throw new Error( - `openai batch ${batchInfo.id} missing ${remaining.size} embedding responses`, - ); - } }, }); } diff --git a/extensions/voyage/embedding-batch.test.ts b/extensions/voyage/embedding-batch.test.ts index 347e46522e4f..43b18e19acaf 100644 --- a/extensions/voyage/embedding-batch.test.ts +++ b/extensions/voyage/embedding-batch.test.ts @@ -135,49 +135,91 @@ afterEach(() => { }); describe("voyage batch bounded reads", () => { - it("preserves configured query parameters on real direct embedding requests", async () => { - const received: Array<{ url: string; authorization: string | undefined }> = []; - const server = createServer((request, response) => { - received.push({ url: request.url ?? "", authorization: request.headers.authorization }); - if (request.url !== "/tenant/v1/embeddings?api-version=2024-10-21&tenant=beta") { - response.writeHead(404).end("wrong embedding endpoint"); - return; + it.each([ + { operation: "single", inputType: "query", expectedInputs: [["first"]] }, + { operation: "single", inputType: "document", expectedInputs: [["first"]] }, + { operation: "single", inputType: undefined, expectedInputs: [["first"]] }, + { operation: "batch", inputType: "query", expectedInputs: [["first"], ["second"]] }, + { operation: "batch", inputType: "document", expectedInputs: [["first", "second"]] }, + { operation: "batch", inputType: undefined, expectedInputs: [["first", "second"]] }, + ] as const)( + "preserves real $operation $inputType requests, grouping, and configured query parameters", + async ({ operation, inputType, expectedInputs }) => { + const received: Array<{ + url: string; + authorization: string | undefined; + body: unknown; + }> = []; + const vectors: Record = { first: [7, 11], second: [13, 17] }; + const server = createServer((request, response) => { + request.setEncoding("utf8"); + let text = ""; + request.on("data", (chunk: string) => { + text += chunk; + }); + request.on("end", () => { + if (request.url !== "/tenant/v1/embeddings?api-version=2024-10-21&tenant=beta") { + response.writeHead(404).end("wrong embedding endpoint"); + return; + } + const body = JSON.parse(text) as { input: string[] }; + received.push({ url: request.url, authorization: request.headers.authorization, body }); + response.writeHead(200, { "content-type": "application/json" }); + response.end( + JSON.stringify({ + data: body.input.map((input, index) => ({ index, embedding: vectors[input] })), + }), + ); + }); + }); + server.listen(0, "127.0.0.1"); + await once(server, "listening"); + const address = server.address(); + if (!address || typeof address === "string") { + throw new Error("expected loopback TCP address"); } - response.writeHead(200, { "content-type": "application/json" }); - response.end(JSON.stringify({ data: [{ embedding: [7, 11] }] })); - }); - server.listen(0, "127.0.0.1"); - await once(server, "listening"); - const address = server.address(); - if (!address || typeof address === "string") { - throw new Error("expected loopback TCP address"); - } - try { - const { provider } = await createVoyageEmbeddingProvider({ - config: {}, - provider: "voyage", - model: "voyage-3", - fallback: "none", - remote: { - baseUrl: `http://127.0.0.1:${address.port}/tenant/v1/?api-version=2024-10-21&tenant=beta#local`, - apiKey: "voyage-loopback-key", - }, - }); - - await expect(provider.embed("hello", { inputType: "query" })).resolves.toEqual([7, 11]); - expect(received).toEqual([ - { - url: "/tenant/v1/embeddings?api-version=2024-10-21&tenant=beta", - authorization: "Bearer voyage-loopback-key", - }, - ]); - } finally { - await new Promise((resolve, reject) => { - server.close((error) => (error ? reject(error) : resolve())); - }); - } - }); + try { + const { provider } = await createVoyageEmbeddingProvider({ + config: {}, + provider: "voyage", + model: "voyage/voyage-3", + fallback: "none", + remote: { + baseUrl: `http://127.0.0.1:${address.port}/tenant/v1/?api-version=2024-10-21&tenant=beta#local`, + apiKey: "voyage-loopback-key", + }, + }); + expect(provider.maxInputTokens).toBe(32000); + if (operation === "single") { + await expect(provider.embed({ text: "first" }, { inputType })).resolves.toEqual([7, 11]); + } else { + await expect( + provider.embedBatch([{ text: "first" }, "second"], { inputType }), + ).resolves.toEqual([ + [7, 11], + [13, 17], + ]); + } + expect(received).toHaveLength(expectedInputs.length); + expect(received).toEqual( + expect.arrayContaining( + expectedInputs.map((input) => ({ + url: "/tenant/v1/embeddings?api-version=2024-10-21&tenant=beta", + authorization: "Bearer voyage-loopback-key", + body: { model: "voyage-3", input, input_type: inputType ?? "document" }, + })), + ), + ); + await expect(provider.embedBatch([], { inputType })).resolves.toEqual([]); + expect(received).toHaveLength(expectedInputs.length); + } finally { + await new Promise((resolve, reject) => { + server.close((error) => (error ? reject(error) : resolve())); + }); + } + }, + ); it("preserves configured query parameters through real batch upload, create, status, and error output", async () => { const received: Array<{ url: string; authorization: string | undefined }> = []; diff --git a/extensions/voyage/embedding-batch.ts b/extensions/voyage/embedding-batch.ts index 7e3781edd4a9..a6c14999b586 100644 --- a/extensions/voyage/embedding-batch.ts +++ b/extensions/voyage/embedding-batch.ts @@ -1,6 +1,5 @@ // Voyage plugin module implements embedding batch behavior. import { - applyEmbeddingBatchOutputLine, buildBatchHeaders, buildEmbeddingBatchGroupOptions, EMBEDDING_BATCH_ENDPOINT, @@ -8,11 +7,8 @@ import { formatBatchErrorDetail, formatUnavailableBatchError, postJsonWithRetry, - readEmbeddingBatchJsonl, resolveEmbeddingEndpointUrl, - resolveCompletedBatchResult, - runEmbeddingBatchGroups, - throwIfBatchCompletionError, + runEmbeddingBatches, type EmbeddingBatchExecutionParams, type EmbeddingBatchStatus, type ProviderBatchOutputLine, @@ -160,106 +156,50 @@ export async function runVoyageEmbeddingBatches( requests: VoyageBatchRequest[]; } & EmbeddingBatchExecutionParams, ): Promise> { - return await runEmbeddingBatchGroups({ + return await runEmbeddingBatches({ + provider: "voyage", ...buildEmbeddingBatchGroupOptions(params, { maxRequests: VOYAGE_BATCH_MAX_REQUESTS, debugLabel: "memory embeddings: voyage batch submit", }), - runGroup: async ({ group, groupIndex, groups, byCustomId, pollIntervalMs, timeoutMs }) => { - const batchInfo = await submitVoyageBatch({ - client: params.client, - requests: group, - agentId: params.agentId, - }); - if (!batchInfo.id) { - throw new Error("voyage batch create failed: missing batch id"); - } + submit: (group) => + submitVoyageBatch({ client: params.client, requests: group, agentId: params.agentId }), + readError: (errorFileId) => readVoyageBatchError({ client: params.client, errorFileId }), + readOutput: (fileId, read) => + withRemoteHttpResponse( + buildVoyageBatchRequest({ + client: params.client, + path: `files/${fileId}/content`, + onResponse: async (response) => { + await assertOkOrThrowProviderError(response, "voyage.batch-file-content"); + await read(response); + }, + }), + ), + waitForBatch: async (batchInfo, pollIntervalMs, timeoutMs) => { const batchId = batchInfo.id; - - params.debug?.("memory embeddings: voyage batch created", { - batchId: batchInfo.id, - status: batchInfo.status, - group: groupIndex + 1, - groups, - requests: group.length, + const client = params.client; + const wait = params.wait; + const debug = params.debug; + const deadline = createProviderOperationDeadline({ + label: `voyage batch ${batchId}`, + timeoutMs, }); - - await throwIfBatchCompletionError({ + return await waitForEmbeddingBatch({ provider: "voyage", - status: batchInfo, - readError: async (errorFileId) => - await readVoyageBatchError({ client: params.client, errorFileId }), + batchId, + wait, + pollIntervalMs, + timeoutMs, + debug, + initial: batchInfo, + fetchStatus: (signal) => fetchVoyageBatchStatus({ client, batchId, signal }), + resolveTimeoutMs: () => + resolveProviderOperationTimeoutMs({ deadline, defaultTimeoutMs: timeoutMs }), + waitForPoll: (delayMs) => + waitProviderOperationPollInterval({ deadline, pollIntervalMs: delayMs }), + readError: async (errorFileId) => await readVoyageBatchError({ client, errorFileId }), }); - - const completed = await resolveCompletedBatchResult({ - provider: "voyage", - status: batchInfo, - wait: params.wait, - waitForBatch: async () => { - const client = params.client; - const wait = params.wait; - const debug = params.debug; - const deadline = createProviderOperationDeadline({ - label: `voyage batch ${batchId}`, - timeoutMs, - }); - return await waitForEmbeddingBatch({ - provider: "voyage", - batchId, - wait, - pollIntervalMs, - timeoutMs, - debug, - initial: batchInfo, - fetchStatus: (signal) => fetchVoyageBatchStatus({ client, batchId, signal }), - resolveTimeoutMs: () => - resolveProviderOperationTimeoutMs({ deadline, defaultTimeoutMs: timeoutMs }), - waitForPoll: (delayMs) => - waitProviderOperationPollInterval({ deadline, pollIntervalMs: delayMs }), - readError: async (errorFileId) => await readVoyageBatchError({ client, errorFileId }), - }); - }, - }); - - const errors: string[] = []; - const remaining = new Set(group.map((request) => request.custom_id)); - - await withRemoteHttpResponse({ - url: resolveEmbeddingEndpointUrl( - params.client.baseUrl, - `files/${completed.outputFileId}/content`, - ), - ssrfPolicy: params.client.ssrfPolicy, - init: { - headers: buildBatchHeaders(params.client, { json: true }), - }, - onResponse: async (contentRes) => { - await assertOkOrThrowProviderError(contentRes, "voyage.batch-file-content"); - - await readEmbeddingBatchJsonl(contentRes, { - label: "voyage.batch-file-content", - maxRecords: group.length, - onRecord: (line) => { - // Only the first response for a submitted id may mutate results. - if (line.custom_id && remaining.has(line.custom_id)) { - applyEmbeddingBatchOutputLine({ line, remaining, errors, byCustomId }); - } - return errors.length === 0 && remaining.size > 0; - }, - }); - }, - }); - - if (errors.length > 0) { - throw new Error( - `voyage batch ${batchInfo.id} failed: ${formatBatchErrorDetail(errors[0]) ?? "unknown error"}`, - ); - } - if (remaining.size > 0) { - throw new Error( - `voyage batch ${batchInfo.id} missing ${remaining.size} embedding responses`, - ); - } }, }); } diff --git a/extensions/voyage/embedding-provider.ts b/extensions/voyage/embedding-provider.ts index 79d06df5ae29..ced7bd8c6e98 100644 --- a/extensions/voyage/embedding-provider.ts +++ b/extensions/voyage/embedding-provider.ts @@ -1,9 +1,8 @@ // Voyage provider module implements model/runtime integration. import { - fetchRemoteEmbeddingVectors, + createRemoteEmbeddingProvider, normalizeEmbeddingModelWithPrefixes, - resolveEmbeddingEndpointUrl, - resolveRemoteEmbeddingBearerClient, + resolveRemoteEmbeddingClient, type MemoryEmbeddingProvider, type MemoryEmbeddingProviderCreateOptions, } from "openclaw/plugin-sdk/memory-core-host-engine-embeddings"; @@ -35,74 +34,18 @@ function normalizeVoyageModel(model: string): string { export async function createVoyageEmbeddingProvider( options: MemoryEmbeddingProviderCreateOptions, ): Promise<{ provider: MemoryEmbeddingProvider; client: VoyageEmbeddingClient }> { - const client = await resolveVoyageEmbeddingClient(options); - const url = resolveEmbeddingEndpointUrl(client.baseUrl, "embeddings"); - - const embedMany = async ( - input: string[], - input_type?: "query" | "document", - signal?: AbortSignal, - ): Promise => { - if (input.length === 0) { - return []; - } - const body: { model: string; input: string[]; input_type?: "query" | "document" } = { - model: client.model, - input, - }; - if (input_type) { - body.input_type = input_type; - } - - return await fetchRemoteEmbeddingVectors({ - url, - headers: client.headers, - ssrfPolicy: client.ssrfPolicy, - signal, - body, - errorPrefix: "voyage embeddings failed", - }); - }; - - return { - provider: { - id: "voyage", - model: client.model, - maxInputTokens: VOYAGE_MAX_INPUT_TOKENS[client.model], - embed: async (input, optionsValue) => { - const text = typeof input === "string" ? input : input.text; - const [vec] = await embedMany( - [text], - optionsValue?.inputType === "query" ? "query" : "document", - optionsValue?.signal, - ); - return vec ?? []; - }, - embedBatch: async (inputs, optionsLocal) => { - const texts = inputs.map((input) => (typeof input === "string" ? input : input.text)); - if (optionsLocal?.inputType === "query") { - return await Promise.all( - texts.map(async (text) => { - const [vec] = await embedMany([text], "query", optionsLocal.signal); - return vec ?? []; - }), - ); - } - return await embedMany(texts, "document", optionsLocal?.signal); - }, - }, - client, - }; -} - -async function resolveVoyageEmbeddingClient( - options: MemoryEmbeddingProviderCreateOptions, -): Promise { - const { baseUrl, headers, ssrfPolicy } = await resolveRemoteEmbeddingBearerClient({ + const client = await resolveRemoteEmbeddingClient({ provider: "voyage", options, defaultBaseUrl: DEFAULT_VOYAGE_BASE_URL, + normalizeModel: normalizeVoyageModel, }); - const model = normalizeVoyageModel(options.model); - return { baseUrl, headers, ssrfPolicy, model }; + const provider = createRemoteEmbeddingProvider({ + id: "voyage", + client, + errorPrefix: "voyage embeddings failed", + buildRequestFields: (kind) => ({ input_type: kind }), + }); + provider.maxInputTokens = VOYAGE_MAX_INPUT_TOKENS[client.model]; + return { provider, client }; } diff --git a/packages/memory-host-sdk/src/engine-embeddings.ts b/packages/memory-host-sdk/src/engine-embeddings.ts index 6868deb525d5..93a79f490ec4 100644 --- a/packages/memory-host-sdk/src/engine-embeddings.ts +++ b/packages/memory-host-sdk/src/engine-embeddings.ts @@ -32,6 +32,7 @@ export { export { buildEmbeddingBatchGroupOptions, runEmbeddingBatchGroups, + runEmbeddingBatches, type EmbeddingBatchExecutionParams, } from "./host/batch-runner.js"; export { diff --git a/packages/memory-host-sdk/src/host/batch-runner.ts b/packages/memory-host-sdk/src/host/batch-runner.ts index dbfa84416282..9c6ed4297351 100644 --- a/packages/memory-host-sdk/src/host/batch-runner.ts +++ b/packages/memory-host-sdk/src/host/batch-runner.ts @@ -1,5 +1,13 @@ // Memory Host SDK module implements batch runner behavior. import { resolveSafeTimeoutDelayMs } from "../../../gateway-client/src/timeouts.js"; +import { formatBatchErrorDetail } from "./batch-error-utils.js"; +import { applyEmbeddingBatchOutputLine, readEmbeddingBatchJsonl } from "./batch-output.js"; +import type { EmbeddingBatchStatus, ProviderBatchOutputLine } from "./batch-provider-common.js"; +import { + resolveCompletedBatchResult, + throwIfBatchCompletionError, + type BatchCompletionResult, +} from "./batch-status.js"; import { splitBatchRequestsByLimits } from "./batch-utils.js"; import { runMemoryHostTasksWithConcurrency } from "./internal.js"; @@ -139,3 +147,77 @@ export function buildEmbeddingBatchGroupOptions( debugLabel: options.debugLabel, }; } + +/** Run compatible batch jobs while providers retain submission, polling, and HTTP ownership. */ +export async function runEmbeddingBatches< + TRequest extends { custom_id: string }, + TStatus extends EmbeddingBatchStatus, +>( + params: Omit>[0], "runGroup"> & { + provider: string; + submit: (group: TRequest[]) => Promise; + waitForBatch: ( + status: TStatus & { id: string }, + pollIntervalMs: number, + timeoutMs: number, + ) => Promise; + readError: (errorFileId: string) => Promise; + readOutput: (fileId: string, read: (response: Response) => Promise) => Promise; + }, +): Promise> { + return await runEmbeddingBatchGroups({ + ...params, + runGroup: async ({ group, groupIndex, groups, byCustomId, pollIntervalMs, timeoutMs }) => { + const status = await params.submit(group); + if (!status.id) { + throw new Error(`${params.provider} batch create failed: missing batch id`); + } + const batchId = status.id; + params.debug?.(`memory embeddings: ${params.provider} batch created`, { + batchId, + status: status.status, + group: groupIndex + 1, + groups, + requests: group.length, + }); + // A completed error file takes precedence over requiring or downloading success output. + await throwIfBatchCompletionError({ + provider: params.provider, + status, + readError: params.readError, + }); + const completed = await resolveCompletedBatchResult({ + provider: params.provider, + status, + wait: params.wait, + waitForBatch: () => + params.waitForBatch({ ...status, id: batchId }, pollIntervalMs, timeoutMs), + }); + const errors: string[] = []; + const remaining = new Set(group.map((request) => request.custom_id)); + await params.readOutput(completed.outputFileId, async (response) => { + await readEmbeddingBatchJsonl(response, { + label: `${params.provider}.batch-file-content`, + maxRecords: group.length, + onRecord: (line) => { + // Only the first response for a submitted id may mutate results. + if (line.custom_id && remaining.has(line.custom_id)) { + applyEmbeddingBatchOutputLine({ line, remaining, errors, byCustomId }); + } + return errors.length === 0 && remaining.size > 0; + }, + }); + }); + if (errors.length > 0) { + throw new Error( + `${params.provider} batch ${batchId} failed: ${formatBatchErrorDetail(errors[0]) ?? "unknown error"}`, + ); + } + if (remaining.size > 0) { + throw new Error( + `${params.provider} batch ${batchId} missing ${remaining.size} embedding responses`, + ); + } + }, + }); +} diff --git a/src/plugin-sdk/memory-core-host-engine-embeddings.ts b/src/plugin-sdk/memory-core-host-engine-embeddings.ts index 674431affb63..b7f2affc9eac 100644 --- a/src/plugin-sdk/memory-core-host-engine-embeddings.ts +++ b/src/plugin-sdk/memory-core-host-engine-embeddings.ts @@ -47,6 +47,7 @@ export { resolveRemoteEmbeddingBearerClient, resolveRemoteEmbeddingClient, runEmbeddingBatchGroups, + runEmbeddingBatches, sanitizeAndNormalizeEmbedding, sanitizeEmbeddingCacheHeaders, throwIfBatchCompletionError,