From c85a5c57bafce1078737a8938d10826a7737ac1c Mon Sep 17 00:00:00 2001 From: zie Date: Thu, 23 Jul 2026 16:06:16 +0700 Subject: [PATCH] fix(usage): record exact embedding tokens --- open-sse/handlers/embeddingsCore.js | 1 + src/sse/handlers/embeddings.js | 26 +++++- .../unit/embedding-usage-persistence.test.js | 91 +++++++++++++++++++ 3 files changed, 117 insertions(+), 1 deletion(-) create mode 100644 tests/unit/embedding-usage-persistence.test.js diff --git a/open-sse/handlers/embeddingsCore.js b/open-sse/handlers/embeddingsCore.js index 5a4c92ba..aa81117c 100644 --- a/open-sse/handlers/embeddingsCore.js +++ b/open-sse/handlers/embeddingsCore.js @@ -116,6 +116,7 @@ export async function handleEmbeddingsCore({ return { success: true, + usage: normalized.usage || null, response: new Response(JSON.stringify(normalized), { headers: { "Content-Type": "application/json", diff --git a/src/sse/handlers/embeddings.js b/src/sse/handlers/embeddings.js index cde0d41e..45e4d171 100644 --- a/src/sse/handlers/embeddings.js +++ b/src/sse/handlers/embeddings.js @@ -12,6 +12,16 @@ import { errorResponse, unavailableResponse } from "open-sse/utils/error.js"; import { HTTP_STATUS } from "open-sse/config/runtimeConfig.js"; import * as log from "../utils/logger.js"; import { updateProviderCredentials, checkAndRefreshToken } from "../services/tokenRefresh.js"; +import { saveRequestUsage } from "@/lib/usageDb.js"; + +function exactEmbeddingUsage(raw) { + if (!raw || typeof raw !== "object" || Array.isArray(raw) || raw.estimated === true) return null; + const promptTokens = raw.prompt_tokens ?? raw.input_tokens; + const completionTokens = raw.completion_tokens ?? raw.output_tokens ?? 0; + const totalTokens = raw.total_tokens; + if (!Number.isSafeInteger(promptTokens) || promptTokens <= 0 || completionTokens !== 0 || totalTokens !== promptTokens) return null; + return { prompt_tokens: promptTokens, completion_tokens: 0, total_tokens: totalTokens }; +} /** * Handle embeddings request for the SSE/Next.js server. @@ -124,7 +134,21 @@ export async function handleEmbeddings(request) { } }); - if (result.success) return result.response; + if (result.success) { + const usage = exactEmbeddingUsage(result.usage); + if (usage) { + saveRequestUsage({ + provider, + model, + connectionId: credentials.connectionId, + apiKey, + endpoint: url.pathname, + tokens: usage, + status: "success", + }).catch(() => {}); + } + return result.response; + } const { shouldFallback } = await markAccountUnavailable(credentials.connectionId, result.status, result.error, provider, model); diff --git a/tests/unit/embedding-usage-persistence.test.js b/tests/unit/embedding-usage-persistence.test.js new file mode 100644 index 00000000..9fcef0d7 --- /dev/null +++ b/tests/unit/embedding-usage-persistence.test.js @@ -0,0 +1,91 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const mocks = vi.hoisted(() => ({ + handleEmbeddingsCore: vi.fn(), + saveRequestUsage: vi.fn(), +})); + +vi.mock("../../src/sse/services/auth.js", () => ({ + getProviderCredentials: async () => ({ + apiKey: "provider-secret", + connectionId: "connection-a", + connectionName: "Provider A", + }), + markAccountUnavailable: vi.fn(), + clearAccountError: vi.fn(), + extractApiKey: () => "client-key", + isValidApiKey: vi.fn(), +})); +vi.mock("@/lib/localDb", () => ({ getSettings: async () => ({ requireApiKey: false }) })); +vi.mock("../../src/sse/services/model.js", () => ({ + getModelInfo: async () => ({ provider: "openai", model: "text-embedding-3-small" }), +})); +vi.mock("../../open-sse/handlers/embeddingsCore.js", () => ({ + handleEmbeddingsCore: mocks.handleEmbeddingsCore, +})); +vi.mock("../../open-sse/utils/error.js", () => ({ + errorResponse: (status, message) => Response.json({ error: message }, { status }), + unavailableResponse: (status, message) => Response.json({ error: message }, { status }), +})); +vi.mock("../../src/sse/utils/logger.js", () => ({ + request: vi.fn(), debug: vi.fn(), warn: vi.fn(), error: vi.fn(), info: vi.fn(), maskKey: vi.fn(), +})); +vi.mock("../../src/sse/services/tokenRefresh.js", () => ({ + updateProviderCredentials: vi.fn(), + checkAndRefreshToken: async (_provider, credentials) => credentials, +})); +vi.mock("@/lib/usageDb.js", () => ({ saveRequestUsage: mocks.saveRequestUsage })); + +import { handleEmbeddings } from "../../src/sse/handlers/embeddings.js"; + +describe("embedding usage persistence", () => { + beforeEach(() => { + vi.clearAllMocks(); + mocks.saveRequestUsage.mockResolvedValue(undefined); + mocks.handleEmbeddingsCore.mockResolvedValue({ + success: true, + usage: { prompt_tokens: 12, total_tokens: 12 }, + response: Response.json({ data: [] }), + }); + }); + + it("records exact provider usage for successful embedding requests", async () => { + await handleEmbeddings(new Request("http://localhost/v1/embeddings", { + method: "POST", + body: JSON.stringify({ model: "openai/text-embedding-3-small", input: "hello" }), + })); + + expect(mocks.saveRequestUsage).toHaveBeenCalledWith(expect.objectContaining({ + provider: "openai", + model: "text-embedding-3-small", + connectionId: "connection-a", + apiKey: "client-key", + endpoint: "/v1/embeddings", + status: "success", + tokens: { prompt_tokens: 12, completion_tokens: 0, total_tokens: 12 }, + })); + }); + + it.each([ + null, + {}, + { prompt_tokens: 0, total_tokens: 0 }, + { prompt_tokens: "12", total_tokens: 12 }, + { prompt_tokens: 12, total_tokens: 13 }, + { prompt_tokens: 12, completion_tokens: 1, total_tokens: 12 }, + { prompt_tokens: 12, total_tokens: 12, estimated: true }, + ])("does not record inexact usage %#", async (usage) => { + mocks.handleEmbeddingsCore.mockResolvedValue({ + success: true, + usage, + response: Response.json({ data: [] }), + }); + + await handleEmbeddings(new Request("http://localhost/v1/embeddings", { + method: "POST", + body: JSON.stringify({ model: "openai/text-embedding-3-small", input: "hello" }), + })); + + expect(mocks.saveRequestUsage).not.toHaveBeenCalled(); + }); +});