diff --git a/.env.example b/.env.example index a85ba4c9234..569027064cb 100644 --- a/.env.example +++ b/.env.example @@ -1403,6 +1403,20 @@ OPENWEATHER_API_KEY= # or # COHERE_API_KEY=your_cohere_api_key +#======================# +# Classification # +#======================# + +# Key for the provider named by `classification.provider` in librechat.yaml. +# Each provider declares which variable it reads through +# `classification.providers..apiKeyEnv`; this is the default. +# CLASSIFIER_API_KEY=your_classifier_api_key + +# The presets read these instead. +# TYPESAFE_API_KEY=your_typesafe_api_key +# OPENROUTER_KEY=your_openrouter_key +# CLOUDFLARE_API_TOKEN=your_cloudflare_api_token + #======================# # MCP Configuration # #======================# diff --git a/api/server/controllers/agents/client.js b/api/server/controllers/agents/client.js index fcf65b8c028..2360dd29da3 100644 --- a/api/server/controllers/agents/client.js +++ b/api/server/controllers/agents/client.js @@ -45,6 +45,9 @@ const { computeAgentRequestFingerprint, computeLegacyAgentRequestFingerprint, getRunDiscoveredTools, + predictToolsForTurn, + createMemoryGate, + classificationCapability, captureResumeModelParameters, pickResumeContext, getApprovalTtlMs, @@ -1852,6 +1855,40 @@ class AgentClient extends BaseClient { return wiring; } + /** + * Builds the classifier gate that skips the memory model on turns with nothing durable. + * @returns {import('@librechat/api').MemoryGate | null} + */ + buildMemoryGate() { + const config = this.options.req?.config?.classification; + const capability = classificationCapability(config, 'memoryGate'); + if (capability == null) { + return null; + } + return createMemoryGate({ + classifier: capability.classifier, + settings: capability.settings, + onUsage: (usage, model) => this.recordClassifierUsage(usage, model), + }); + } + + /** Bills a classifier judgment like any other secondary call: priced from the shared table. */ + recordClassifierUsage(usage, model) { + const appConfig = this.options.req?.config; + return this.recordCollectedUsage({ + collectedUsage: [{ input_tokens: usage.inputTokens, output_tokens: usage.outputTokens }], + context: 'classification', + model, + crossEndpoint: true, + balance: getBalanceConfig(appConfig), + transactions: getTransactionsConfig(appConfig), + messageId: this.responseMessageId, + updateStreamUsage: false, + }).catch((err) => { + logger.error('[AgentClient] Error recording classifier usage', getSafeErrorMetadata(err)); + }); + } + /** Builds the independently opt-in live reasoning-label controller. */ buildReasoningLabelWiring(streamId, abortSignal, seedFromContent = false) { if (!streamId || typeof Run?.prototype?.generateReasoningLabel !== 'function') { @@ -3383,6 +3420,8 @@ class AgentClient extends BaseClient { res: this.options.res, user: createSafeUser(this.options.req.user), tenantId: resolveRequestTenantId(this.options.req), + /** Null unless `classification.memoryGate` is on. */ + gate: this.buildMemoryGate(), }); this.processMemory = processMemory; @@ -4378,6 +4417,16 @@ class AgentClient extends BaseClient { ); } + /** By the pause these describe what the turn had loaded, so the resumed + * segment must rebuild with them. Run state records real discoveries only. */ + if (this.predictedToolNames?.length) { + const merged = new Set(discoveredTools); + for (const name of this.predictedToolNames) { + merged.add(name); + } + discoveredTools = Array.from(merged); + } + this.stagedApproval = { streamId, pendingAction, @@ -4798,6 +4847,21 @@ class AgentClient extends BaseClient { if (this.agentConfigs && this.agentConfigs.size > 0) { agents.push(...this.agentConfigs.values()); } + + /** Awaited after the memory run starts, so the two overlap. */ + const predictionPromise = predictToolsForTurn({ + config: appConfig?.classification, + agents, + messages, + signal: abortController.signal, + onUsage: (usage, model) => this.recordClassifierUsage(usage, model), + }).catch((error) => { + logger.warn( + '[AgentClient] Tool prediction failed; continuing without it', + getSafeErrorMetadata(error), + ); + return []; + }); const modelBoundCallback = AgentClient.prototype.createModelBoundChatModelCallback.call(this); const initialModelBoundAdmission = @@ -4840,6 +4904,11 @@ class AgentClient extends BaseClient { memoryPromise = this.runMemory(memoryMessages); } + const predictedToolNames = await predictionPromise; + if (predictedToolNames.length > 0) { + this.predictedToolNames = predictedToolNames; + } + const { calibrationRatio, fadingTier, fadingTiers } = resolveRunSeeds(this); const streamId = this.options.req?._resumableStreamId; @@ -4916,6 +4985,7 @@ class AgentClient extends BaseClient { messages, discoveredToolNames: this.eventActorContinuation === 'warm' ? this.eventActorDiscoveredToolNames : undefined, + predictedToolNames, modelCallbacks: [ modelBoundCallback, createAgentMemoryCallback(this.attachmentMemoryContext ?? {}), diff --git a/api/server/controllers/agents/client.test.js b/api/server/controllers/agents/client.test.js index c646df0436d..a32b867ca32 100644 --- a/api/server/controllers/agents/client.test.js +++ b/api/server/controllers/agents/client.test.js @@ -7,6 +7,9 @@ const mockDetachedUsageRecorder = jest.fn(); const mockCreateDetachedSubagentUsageRecorder = jest.fn(() => mockDetachedUsageRecorder); const mockGetAgentCheckpointer = jest.fn(); const mockHasDurableAgentInterruptCheckpoint = jest.fn().mockResolvedValue(true); +const mockPredictToolsForTurn = jest.fn((...args) => + jest.requireActual('@librechat/api').predictToolsForTurn(...args), +); const mockBuildAgentScopedContext = jest.fn((...args) => jest.requireActual('@librechat/api').buildAgentScopedContext(...args), ); @@ -190,6 +193,7 @@ jest.mock('@librechat/api', () => ({ buildAgentScopedContext: (...args) => mockBuildAgentScopedContext(...args), checkAccess: jest.fn(), createRun: (...args) => mockCreateRun(...args), + predictToolsForTurn: (...args) => mockPredictToolsForTurn(...args), countFormattedMessageTokens: jest.fn(() => 42), countTokens: jest.fn((text) => Math.ceil(String(text ?? '').length / 4)), createCachedTokenCounter: jest.fn(async () => jest.fn(() => 0)), @@ -2377,6 +2381,95 @@ describe('AgentClient - startup telemetry', () => { }, ); + it('bills classifier usage as its own secondary transaction', async () => { + mockRecordCollectedUsage.mockClear(); + const client = new AgentClient({ + req: { user: { id: 'user-123' }, body: {}, config: {} }, + res: {}, + agent: { + id: 'agent-123', + endpoint: EModelEndpoint.openAI, + provider: EModelEndpoint.openAI, + model_parameters: { model: 'gpt-4' }, + }, + endpointTokenConfig: {}, + eventHandlers: {}, + contentParts: [], + collectedUsage: [], + artifactPromises: [], + }); + client.responseMessageId = 'classifier-usage-response'; + mockRecordCollectedUsage.mockResolvedValueOnce({ input_tokens: 512, output_tokens: 0 }); + + await client.recordClassifierUsage({ inputTokens: 512, outputTokens: 0 }, 'jev-latest'); + + expect(mockRecordCollectedUsage).toHaveBeenCalledTimes(1); + const [, params] = mockRecordCollectedUsage.mock.calls[0]; + expect(params).toEqual( + expect.objectContaining({ + context: 'classification', + model: 'jev-latest', + messageId: 'classifier-usage-response', + collectedUsage: [{ input_tokens: 512, output_tokens: 0 }], + }), + ); + expect(client.getStreamUsage()).toBeFalsy(); + }); + + it('starts the memory run while tool prediction is still running', async () => { + let resolvePrediction; + mockPredictToolsForTurn.mockImplementationOnce( + () => new Promise((resolve) => (resolvePrediction = resolve)), + ); + mockCreateRun.mockResolvedValueOnce({ + Graph: null, + processStream: jest.fn().mockResolvedValue(), + getCalibrationRatio: jest.fn(() => 0), + getInterrupt: jest.fn(() => undefined), + }); + const createRunBefore = mockCreateRun.mock.calls.length; + const client = new AgentClient({ + req: { user: { id: 'user-123' }, body: {}, config: {} }, + res: {}, + agent: { + id: 'agent-123', + endpoint: EModelEndpoint.openAI, + provider: EModelEndpoint.openAI, + model_parameters: { model: 'gpt-4' }, + hide_sequential_outputs: false, + tools: [{ name: 'read_file' }], + }, + endpointTokenConfig: {}, + eventHandlers: {}, + contentParts: [], + collectedUsage: [], + artifactPromises: [], + }); + client.conversationId = 'prediction-overlap'; + client.responseMessageId = 'prediction-overlap-response'; + client.parentMessageId = 'prediction-overlap-parent'; + client.recordCollectedUsage = jest.fn().mockResolvedValue(); + client.processMemory = jest.fn(); + client.runMemory = jest.fn().mockResolvedValue(undefined); + + const completion = client.chatCompletion({ payload: [] }); + for (let i = 0; i < 50 && client.runMemory.mock.calls.length === 0; i++) { + await new Promise((resolve) => setImmediate(resolve)); + } + + expect(resolvePrediction).toBeDefined(); + expect(client.runMemory).toHaveBeenCalledTimes(1); + expect(mockCreateRun).toHaveBeenCalledTimes(createRunBefore); + + resolvePrediction(['search_mcp_docs']); + await completion; + + expect(mockCreateRun).toHaveBeenCalledTimes(createRunBefore + 1); + expect(mockCreateRun.mock.calls[createRunBefore][0].predictedToolNames).toEqual([ + 'search_mcp_docs', + ]); + }); + it('uses request-scoped hook resolution when deciding whether a scheduled run can pause', async () => { mockIsHITLEnabled.mockReturnValue(true); registerToolApprovalHook((context) => diff --git a/librechat.example.yaml b/librechat.example.yaml index 260912aa780..4d508495160 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -515,6 +515,77 @@ actions: # - 'host.docker.internal:8080' # - '127.0.0.1:8080' +# Classification: small typed judgments (yes/no, pick one) that code can branch +# on. Off unless enabled. The key is read from the environment variable named by +# apiKeyEnv, never written here. `timeoutMs` bounds each judgment, retries +# included. Every enabled capability sends conversation text to the configured +# host: tool selection sends the latest request and the deferred tools' +# descriptions, the memory gate the last few messages. +# classification: +# enabled: true +# provider: http +# providers: +# http: +# baseURL: https://classifier.example.com/v1/classify +# apiKeyEnv: CLASSIFIER_API_KEY +# timeoutMs: 4000 +# +# # Known hosts ship as presets, so naming one is usually enough. Cloudflare +# # is the exception: its URL carries your account id. +# provider: cloudflare +# providers: +# cloudflare: +# baseURL: https://api.cloudflare.com/client/v4/accounts//ai/run +# apiKeyEnv: CLOUDFLARE_API_TOKEN +# +# # A host with no preset needs no code either, only its shape. `dialect` +# # picks the wire vocabulary, `requestKey` nests the body, `responseKey` +# # unwraps the reply. +# provider: inhouse +# providers: +# inhouse: +# baseURL: https://classify.internal/v1/run +# model: your-model +# dialect: systemone +# requestKey: input +# responseKey: result +# apiKeyEnv: INHOUSE_CLASSIFIER_KEY +# +# # Surfaces the deferred tools a turn is likely to need, so their schemas ship +# # with the first model call instead of costing a tool_search round trip. +# # Only ever adds: a tool it passes over stays listed by name and one search away. +# toolSelection: +# enabled: true +# shortlist: 5 +# # Overrides the provider timeout. It runs before the first model call, so +# # it adds up to this much latency to a turn; a large catalog may need more. +# timeoutMs: 6000 +# # Below this probability that any tool is needed, only tools the request +# # names outright are surfaced. +# needsToolThreshold: 0.15 +# # An unsure ranking surfaces this many extra rather than fewer. +# lowConfidenceExtra: 3 +# # Replace the wording of either question without touching code. +# instructions: Which tool should the assistant call first? +# guidance: Prefer the tool whose purpose matches the request. +# +# # Skips the memory model on turns that do not ask to remember, update or +# # forget anything, which is all the default memory instructions act on. +# memoryGate: +# enabled: true +# threshold: 0.5 +# # If your memory.instructions also keep facts the user did not ask to +# # keep, ask about those instead. whenTrue/whenFalse rather than +# # true/false: YAML reads those bare keys as booleans. +# instructions: Does `latest` say something about the user worth keeping? +# whenTrue: A lasting preference or a fact about the user. +# whenFalse: Small talk, or a detail that only matters in this task. +# # Also ask which of memory.validKeys the turn belongs under and suggest it +# # to the memory model. Rides in the request the gate already makes. +# categorize: true +# categoryThreshold: 0.4 +# detectUpdates: true + # Example MCP Servers Object Structure # mcpServers: # everything: diff --git a/packages/api/src/agents/memory.spec.ts b/packages/api/src/agents/memory.spec.ts index 918f9797821..c7fda4c0601 100644 --- a/packages/api/src/agents/memory.spec.ts +++ b/packages/api/src/agents/memory.spec.ts @@ -1,4 +1,5 @@ import { Types } from 'mongoose'; +import { classificationSchema } from 'librechat-data-provider'; import { Run, Providers, GraphEvents } from '@librechat/agents'; import { AIMessage, HumanMessage } from '@librechat/agents/langchain/messages'; import { Tools, MemoryScope, EModelEndpoint, AgentCapabilities } from 'librechat-data-provider'; @@ -6,6 +7,7 @@ import type { FiltersConfig } from 'librechat-data-provider'; import type { RuntimeProviderName } from '@librechat/agents'; import type { IUser } from '@librechat/data-schemas'; import type { Response } from 'express'; +import type { ClassificationRequest, Classifier } from '~/classification'; import type { ServerRequest } from '~/types'; import { processMemory, @@ -20,6 +22,7 @@ import { buildInlineMemoryContext, } from './memory'; import { GenerationJobManager } from '~/stream/GenerationJobManager'; +import { createMemoryGate } from '~/memory/gate'; jest.mock('~/middleware/access', () => ({ checkAccess: jest.fn().mockResolvedValue(true), @@ -776,6 +779,76 @@ describe('createMemoryTool tokenLimit enforcement', () => { }); }); +describe('memory gate inside the memory processor', () => { + function recordingGate(probability: number) { + const requests: ClassificationRequest[] = []; + const classifier: Classifier = { + id: 'recording', + model: 'test', + async classify(request) { + requests.push(request); + return { + model: 'test', + answers: { request: { type: 'boolean', probability } }, + usage: { inputTokens: 0, outputTokens: 0 }, + }; + }, + }; + const settings = classificationSchema.parse({ memoryGate: { enabled: true } }).memoryGate; + return { requests, gate: createMemoryGate({ classifier, settings }) ?? undefined }; + } + + async function processorWith(gate: ReturnType['gate']) { + const [, process] = await createMemoryProcessor({ + res: { headersSent: false, write: jest.fn() } as unknown as Response, + userId: 'user-1', + messageId: 'message-1', + conversationId: 'conversation-1', + gate, + memoryMethods: { + setMemory: jest.fn().mockResolvedValue({ ok: true }), + deleteMemory: jest.fn().mockResolvedValue({ ok: true }), + getUserMemories: jest.fn().mockResolvedValue([]), + getFormattedMemories: jest.fn().mockResolvedValue({ + withKeys: '', + withoutKeys: '', + totalTokens: 0, + tokenCountsByKey: new Map(), + }), + }, + }); + return process; + } + + const window = [ + new HumanMessage('find me images of wildlife in nagoya'), + new AIMessage(`page snapshot ${'x'.repeat(40_000)}`), + new HumanMessage('I prefer answers in Japanese from now on'), + ]; + const buffer = [new HumanMessage(window.map((m) => m.content).join('\n'))]; + + it('judges the newest turn even when older tool output fills the buffer', async () => { + const { requests, gate } = recordingGate(0.9); + + await ( + await processorWith(gate) + )(buffer, window); + + const conversation = JSON.stringify(requests[0].state); + expect(conversation).toContain('I prefer answers in Japanese from now on'); + }); + + it('skips the memory model when the gate says the turn asks for nothing', async () => { + const { gate } = recordingGate(0.01); + const runCalls = (Run.create as jest.Mock).mock.calls.length; + + const result = await (await processorWith(gate))(buffer, window); + + expect(result).toBeUndefined(); + expect((Run.create as jest.Mock).mock.calls.length).toBe(runCalls); + }); +}); + describe('memory token limit guidance', () => { it('describes the aggregate limit and never reports negative remaining capacity', async () => { const [, process] = await createMemoryProcessor({ diff --git a/packages/api/src/agents/memory.ts b/packages/api/src/agents/memory.ts index 9e7627d9aa2..590e309f6cb 100644 --- a/packages/api/src/agents/memory.ts +++ b/packages/api/src/agents/memory.ts @@ -36,6 +36,7 @@ import type { BaseMessage, ToolMessage } from '@librechat/agents/langchain/messa import type { DynamicStructuredTool } from '@librechat/agents/langchain/tools'; import type { Response as ServerResponse } from 'express'; import type { ServerRequest, RunLLMConfig } from '~/types'; +import type { MemoryGate } from '~/memory/gate'; import { resolveConfigHeaders, createSafeUser, getSafeErrorMetadata } from '~/utils'; import { contentFilterModelBoundBlockResponse } from '~/middleware/contentFilter'; import { extractMemoryContent } from '~/protection/adapters/submissions'; @@ -1031,6 +1032,7 @@ export async function createMemoryProcessor({ jobCreatedAt, user, tenantId, + gate, }: { res: ServerResponse; messageId: string; @@ -1045,6 +1047,8 @@ export async function createMemoryProcessor({ jobCreatedAt?: number; user?: IUser; tenantId?: string; + /** Injected, so this module needs no knowledge of what does the judging. */ + gate?: MemoryGate; }): Promise< [ string, @@ -1074,6 +1078,22 @@ export async function createMemoryProcessor({ messages: BaseMessage[], inspectionMessages?: BaseMessage[], ): Promise<(TAttachment | null)[] | undefined> { + let turnInstructions = finalInstructions; + if (gate != null) { + /** `messages` is one flattened buffer, oldest text first, so the gate reads the window. */ + const judgment = await gate({ messages: inspectionMessages ?? messages, validKeys }); + if (!judgment.process) { + logger.debug('[MemoryAgent] Turn asks for no memory change; skipping', { + userId, + conversationId, + messageId, + }); + return undefined; + } + if (judgment.hint != null) { + turnInstructions = `${finalInstructions}\n\n${judgment.hint}`; + } + } try { return await processMemory({ res, @@ -1093,7 +1113,7 @@ export async function createMemoryProcessor({ totalTokens: totalTokens || 0, tokenCountsByKey, filters, - instructions: finalInstructions, + instructions: turnInstructions, setMemory: memoryMethods.setMemory, deleteMemory: memoryMethods.deleteMemory, user, diff --git a/packages/api/src/agents/run.ts b/packages/api/src/agents/run.ts index c6287521cc7..e9b91e198e2 100644 --- a/packages/api/src/agents/run.ts +++ b/packages/api/src/agents/run.ts @@ -188,7 +188,7 @@ function parseToolSearchLegacy(content: string, discoveredTools: Set): v * @param messages - The conversation message history * @returns Set of tool names that were discovered via tool_search */ -export function extractDiscoveredToolsFromHistory(messages: BaseMessage[]): Set { +export function extractDiscoveredToolsFromHistory(messages: readonly BaseMessage[]): Set { const discoveredTools = new Set(); for (const message of messages) { @@ -2074,6 +2074,7 @@ export async function createRun({ agents, messages, discoveredToolNames, + predictedToolNames, requestBody, codeApprovalMode: requestedCodeApprovalMode, user, @@ -2142,6 +2143,11 @@ export async function createRun({ * replayed here. Merged with (not replacing) names extracted from `messages`. */ discoveredToolNames?: string[]; + /** + * Separate from `discoveredToolNames` on purpose: that one is persisted and + * replayed on resume, and a guess must not be recorded as a real discovery. + */ + predictedToolNames?: string[]; summarizationConfig?: SummarizationConfig; /** * Manual compaction: the primary agent summarizes the history outright and @@ -2296,6 +2302,11 @@ export async function createRun({ discoveredTools.add(name); } } + if (predictedToolNames?.length) { + for (const name of predictedToolNames) { + discoveredTools.add(name); + } + } } /** Admin kill switch for the ask tool — see {@link isAskUserQuestionAdminDisabled}. */ diff --git a/packages/api/src/classification/config.spec.ts b/packages/api/src/classification/config.spec.ts new file mode 100644 index 00000000000..6b6973ebce2 --- /dev/null +++ b/packages/api/src/classification/config.spec.ts @@ -0,0 +1,135 @@ +import { classificationSchema } from 'librechat-data-provider'; +import type { TClassificationConfig } from 'librechat-data-provider'; +import type { ProviderFetch } from './providers/transport'; +import { resolveClassifier } from './resolve'; +import { boolean } from './questions'; + +/** + * Covers the path an operator's `librechat.yaml` actually travels: the zod + * schema, then the resolver, then the bytes on the wire. A field that parses + * but never reaches the request is the failure this file exists to catch. + */ + +const ANSWER = { model: 'jev-1.13.0', answers: { d: { type: 'noul', noul: 0.8 } }, usage: {} }; + +function recorder(response: unknown = ANSWER) { + const calls: { url: string; body: Record }[] = []; + const fetch: ProviderFetch = async (url, init) => { + calls.push({ url, body: JSON.parse(init.body) }); + return { + ok: true, + status: 200, + headers: { get: () => null }, + text: async () => JSON.stringify(response), + }; + }; + return { calls, fetch }; +} + +function parse(raw: unknown): TClassificationConfig { + return classificationSchema.parse(raw); +} + +describe('classification config', () => { + it('parses an unset block into every capability off', () => { + const config = parse({}); + + expect(config.enabled).toBe(false); + expect(config.provider).toBe('http'); + expect(config.providers).toEqual({}); + }); + + it('keeps the wire-shape fields through parsing', () => { + const config = parse({ + enabled: true, + provider: 'cloudflare', + providers: { + cloudflare: { + baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run', + model: 'typesafe/jev', + dialect: 'systemone', + requestKey: 'input', + responseKey: 'result', + apiKeyEnv: 'CLOUDFLARE_API_TOKEN', + timeoutMs: 9000, + maxRetries: 1, + }, + }, + }); + + expect(config.providers.cloudflare).toEqual({ + baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run', + model: 'typesafe/jev', + dialect: 'systemone', + requestKey: 'input', + responseKey: 'result', + apiKeyEnv: 'CLOUDFLARE_API_TOKEN', + timeoutMs: 9000, + maxRetries: 1, + }); + }); + + it('rejects a dialect it does not implement', () => { + const result = classificationSchema.safeParse({ + enabled: true, + provider: 'x', + providers: { x: { baseURL: 'https://x.test/run', dialect: 'logprobs' } }, + }); + + expect(result.success).toBe(false); + }); + + it('rejects a misspelled key instead of dropping it', () => { + const result = classificationSchema.safeParse({ + enabled: true, + provider: 'x', + providers: { x: { baseURL: 'https://x.test/run', requestkey: 'input' } }, + }); + + expect(result.success).toBe(false); + }); + + it('rejects a baseURL that is not a URL', () => { + const result = classificationSchema.safeParse({ + enabled: true, + provider: 'x', + providers: { x: { baseURL: 'classify.internal' } }, + }); + + expect(result.success).toBe(false); + }); + + it('carries a parsed config all the way onto the wire', async () => { + const config = parse({ + enabled: true, + provider: 'cloudflare', + providers: { + cloudflare: { baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run' }, + }, + }); + const { calls, fetch } = recorder({ result: ANSWER }); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://api.cloudflare.com/client/v4/accounts/abc/ai/run'); + expect(calls[0].body).toHaveProperty('input.questions.d.type', 'noul'); + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('lets a parsed override beat the preset it sits on', async () => { + const config = parse({ + enabled: true, + provider: 'typesafe', + providers: { typesafe: { baseURL: 'https://proxy.internal/systemone', model: 'jev-1.13' } }, + }); + const { calls, fetch } = recorder(); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://proxy.internal/systemone'); + expect(calls[0].body.model).toBe('jev-1.13'); + expect((calls[0].body.questions as Record).d.type).toBe('noul'); + }); +}); diff --git a/packages/api/src/classification/index.ts b/packages/api/src/classification/index.ts new file mode 100644 index 00000000000..fce04f418ea --- /dev/null +++ b/packages/api/src/classification/index.ts @@ -0,0 +1,5 @@ +export * from './types'; +export * from './questions'; +export * from './registry'; +export * from './resolve'; +export type { ProviderFetch, Transport, TransportOptions } from './providers/transport'; diff --git a/packages/api/src/classification/presets.spec.ts b/packages/api/src/classification/presets.spec.ts new file mode 100644 index 00000000000..ba2869ca0c8 --- /dev/null +++ b/packages/api/src/classification/presets.spec.ts @@ -0,0 +1,123 @@ +import { classificationSchema } from 'librechat-data-provider'; +import type { TClassificationConfig, TClassificationProviderConfig } from 'librechat-data-provider'; +import type { ProviderFetch } from './providers/transport'; +import { PRESETS, presetFor, mergeSettings } from './registry'; +import { resolveClassifier } from './resolve'; +import { boolean } from './questions'; + +type Captured = { url: string; body: Record }; + +function recorder(response: unknown) { + const calls: Captured[] = []; + const fetch: ProviderFetch = async (url, init) => { + calls.push({ url, body: JSON.parse(init.body) }); + return { + ok: true, + status: 200, + headers: { get: () => null }, + text: async () => JSON.stringify(response), + }; + }; + return { calls, fetch }; +} + +function configFor( + provider: string, + settings?: TClassificationProviderConfig, +): TClassificationConfig { + return classificationSchema.parse({ + enabled: true, + provider, + providers: settings == null ? {} : { [provider]: settings }, + }); +} + +const NOUL = { model: 'jev-1.13.0', answers: { d: { type: 'noul', noul: 0.8 } }, usage: {} }; + +describe('presets', () => { + it('ships every known host as settings, not as code', () => { + expect(Object.keys(PRESETS).sort()).toEqual(['cloudflare', 'http', 'openrouter', 'typesafe']); + }); + + it('lets an operator override any field of a preset', () => { + const merged = mergeSettings(presetFor('typesafe'), { model: 'jev-1.13', timeoutMs: 9000 }); + + expect(merged.model).toBe('jev-1.13'); + expect(merged.timeoutMs).toBe(9000); + expect(merged.baseURL).toBe('https://api.typesafe.ai/v1/systemone'); + }); + + it('sends the System One vocabulary for the typesafe preset', async () => { + const { calls, fetch } = recorder(NOUL); + const classifier = resolveClassifier({ config: configFor('typesafe'), apiKey: 'k', fetch }); + + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://api.typesafe.ai/v1/systemone'); + expect((calls[0].body.questions as Record).d.type).toBe('noul'); + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('nests the body and unwraps the envelope for the cloudflare preset', async () => { + const { calls, fetch } = recorder({ result: NOUL, success: true }); + const config = configFor('cloudflare', { + baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run', + }); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].body).toHaveProperty('input.questions.d.type', 'noul'); + expect(calls[0].body.model).toBe('typesafe/jev'); + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('still reads a bare body when the host does not wrap it', async () => { + const { fetch } = recorder(NOUL); + const config = configFor('cloudflare', { + baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run', + }); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('sends a flat body for the openrouter preset', async () => { + const { calls, fetch } = recorder(NOUL); + const classifier = resolveClassifier({ config: configFor('openrouter'), apiKey: 'k', fetch }); + + await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://openrouter.ai/api/alpha/decisions'); + expect(calls[0].body).not.toHaveProperty('input'); + expect(calls[0].body.model).toBe('~typesafe/jev-latest'); + }); + + it('builds a host it has never heard of from config alone', async () => { + const { calls, fetch } = recorder({ data: NOUL }); + const config = configFor('somethingnew', { + baseURL: 'https://classify.internal/v1/run', + model: 'house-classifier', + dialect: 'systemone', + requestKey: 'payload', + responseKey: 'data', + }); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://classify.internal/v1/run'); + expect(calls[0].body).toHaveProperty('payload.questions.d.type', 'noul'); + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('stays off for an unknown name with no baseURL to go on', () => { + expect(resolveClassifier({ config: configFor('mystery'), apiKey: 'k' })).toBeNull(); + }); + + it('stays off when a preset needs a baseURL the operator did not give', () => { + expect(resolveClassifier({ config: configFor('cloudflare'), apiKey: 'k' })).toBeNull(); + }); +}); diff --git a/packages/api/src/classification/providers/dialect.spec.ts b/packages/api/src/classification/providers/dialect.spec.ts new file mode 100644 index 00000000000..83cf4d1563d --- /dev/null +++ b/packages/api/src/classification/providers/dialect.spec.ts @@ -0,0 +1,63 @@ +import { toWireQuestion, readAnswer } from './dialect'; +import { boolean, choice } from '../questions'; + +describe('toWireQuestion', () => { + it('keeps the port vocabulary by default', () => { + expect(toWireQuestion(boolean('durable?'), 'port').type).toBe('boolean'); + }); + + it('renames a yes/no question for System One', () => { + expect(toWireQuestion(boolean('durable?'), 'systemone').type).toBe('noul'); + }); + + it('leaves a choice alone in either dialect', () => { + const pick = choice('which', { a: null, b: null }); + + expect(toWireQuestion(pick, 'systemone').type).toBe('choice'); + expect(toWireQuestion(pick, 'port').criteria).toEqual({ a: null, b: null }); + }); + + it('omits criteria when the question carries none', () => { + expect(toWireQuestion(boolean('durable?'), 'port')).not.toHaveProperty('criteria'); + }); +}); + +describe('readAnswer', () => { + it('reads a port boolean', () => { + expect(readAnswer({ type: 'boolean', probability: 0.9 }, 'port')).toEqual({ + type: 'boolean', + probability: 0.9, + }); + }); + + it('reads a System One noul as a boolean', () => { + expect(readAnswer({ type: 'noul', noul: 0.9 }, 'systemone')).toEqual({ + type: 'boolean', + probability: 0.9, + }); + }); + + it('does not read a noul when the port dialect was asked for', () => { + expect(readAnswer({ type: 'noul', noul: 0.9 }, 'port')).toBeNull(); + }); + + it('reports an unmeasured confidence as null rather than zero', () => { + const answer = readAnswer( + { type: 'choice', choice: 'a', probabilities: { a: 1 } }, + 'systemone', + ); + + expect(answer).toEqual({ + type: 'choice', + choice: 'a', + confidence: null, + probabilities: { a: 1 }, + }); + }); + + it('drops an answer it cannot read rather than inventing one', () => { + expect(readAnswer({ type: 'choice' }, 'port')).toBeNull(); + expect(readAnswer(null, 'port')).toBeNull(); + expect(readAnswer('yes', 'port')).toBeNull(); + }); +}); diff --git a/packages/api/src/classification/providers/dialect.ts b/packages/api/src/classification/providers/dialect.ts new file mode 100644 index 00000000000..bcea38561f8 --- /dev/null +++ b/packages/api/src/classification/providers/dialect.ts @@ -0,0 +1,46 @@ +import type { TClassificationProviderConfig } from 'librechat-data-provider'; +import type { ClassificationAnswer, ClassificationQuestion } from '../types'; + +export type Dialect = NonNullable; + +interface WireQuestion { + type: string; + instructions: unknown; + criteria?: unknown; +} + +interface WireAnswer { + type?: unknown; + probability?: unknown; + noul?: unknown; + choice?: unknown; + confidence?: unknown; + probabilities?: unknown; +} + +/** System One calls a yes/no question a `noul`; the port calls it a boolean. */ +export function toWireQuestion(question: ClassificationQuestion, dialect: Dialect): WireQuestion { + const type = dialect === 'systemone' && question.type === 'boolean' ? 'noul' : question.type; + return question.criteria == null + ? { type, instructions: question.instructions } + : { type, instructions: question.instructions, criteria: question.criteria }; +} + +export function readAnswer(answer: unknown, dialect: Dialect): ClassificationAnswer | null { + if (answer == null || typeof answer !== 'object') { + return null; + } + const record = answer as WireAnswer; + const probabilities = (record.probabilities ?? {}) as Record; + const confidence = typeof record.confidence === 'number' ? record.confidence : null; + + const booleanType = dialect === 'systemone' ? 'noul' : 'boolean'; + const probability = dialect === 'systemone' ? record.noul : record.probability; + if (record.type === booleanType && typeof probability === 'number') { + return { type: 'boolean', probability }; + } + if (record.type === 'choice' && typeof record.choice === 'string') { + return { type: 'choice', choice: record.choice, confidence, probabilities }; + } + return null; +} diff --git a/packages/api/src/classification/providers/http.spec.ts b/packages/api/src/classification/providers/http.spec.ts new file mode 100644 index 00000000000..79ce8deb494 --- /dev/null +++ b/packages/api/src/classification/providers/http.spec.ts @@ -0,0 +1,340 @@ +import type { ProviderFetch } from './transport'; +import { ClassificationError } from '../types'; +import { createHttpClassifier } from './http'; + +interface StubResponse { + ok: boolean; + status: number; + body: string; + headers?: Record; +} + +function stubTransport(queue: Array): { + transport: ProviderFetch; + calls: Array<{ url: string; headers: Record; body: string }>; +} { + const calls: Array<{ url: string; headers: Record; body: string }> = []; + let index = 0; + const transport: ProviderFetch = async (url, init) => { + calls.push({ url, headers: init.headers, body: init.body }); + const next = queue[Math.min(index, queue.length - 1)]; + index++; + if (next instanceof Error) { + throw next; + } + return { + ok: next.ok, + status: next.status, + headers: { get: (name: string) => next.headers?.[name.toLowerCase()] ?? null }, + text: async () => next.body, + }; + }; + return { transport, calls }; +} + +const ENDPOINT = 'https://classifier.test/v1/classify'; + +const ANSWER = JSON.stringify({ + model: 'test-1', + answers: { verdict: { type: 'boolean', probability: 0.82 } }, + usage: { input_tokens: 120, output_tokens: 8 }, +}); + +const QUESTION = { + verdict: { type: 'boolean' as const, instructions: 'Is this urgent?' }, +}; + +function build(queue: Array, overrides = {}) { + const { transport, calls } = stubTransport(queue); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + fetch: transport, + sleep: async () => undefined, + ...overrides, + }); + return { classifier, calls }; +} + +describe('createHttpClassifier', () => { + it('posts the state and questions and reads the answers back', async () => { + const { classifier, calls } = build([{ ok: true, status: 200, body: ANSWER }]); + + const result = await classifier.classify({ state: 'payouts are failing', questions: QUESTION }); + + expect(result.answers.verdict).toEqual({ type: 'boolean', probability: 0.82 }); + expect(result.usage).toEqual({ inputTokens: 120, outputTokens: 8 }); + expect(calls[0].url).toBe(ENDPOINT); + expect(calls[0].headers.Authorization).toBe('Bearer sk-test'); + expect(JSON.parse(calls[0].body)).toEqual({ + state: 'payouts are failing', + questions: QUESTION, + }); + }); + + it('includes the model only when one is configured', async () => { + const { classifier, calls } = build([{ ok: true, status: 200, body: ANSWER }], { + model: 'test-1', + }); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(JSON.parse(calls[0].body).model).toBe('test-1'); + expect(classifier.model).toBe('test-1'); + expect(classifier.id).toBe('http'); + }); + + it('carries a choice answer through and drops a type it does not know', async () => { + const body = JSON.stringify({ + model: 'test-1', + answers: { + pick: { type: 'choice', choice: 'b', confidence: 0.7, probabilities: { a: 0.3, b: 0.7 } }, + rate: { type: 'score', score: 1.4, confidence: 0.5, probabilities: { '0': 0.6, '1': 0.4 } }, + }, + usage: { input_tokens: 10, output_tokens: 2 }, + }); + const { classifier } = build([{ ok: true, status: 200, body }]); + + const result = await classifier.classify({ + state: 'x', + questions: { + pick: { type: 'choice', instructions: 'which', criteria: { a: null, b: null } }, + }, + }); + + expect(result.answers.pick).toMatchObject({ type: 'choice', choice: 'b', confidence: 0.7 }); + expect(result.answers).not.toHaveProperty('rate'); + }); + + it('retries a 429 and honors retry-after', async () => { + const waits: number[] = []; + const { transport, calls } = stubTransport([ + { ok: false, status: 429, body: 'slow down', headers: { 'retry-after': '2' } }, + { ok: true, status: 200, body: ANSWER }, + ]); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + fetch: transport, + sleep: async (ms) => { + waits.push(ms); + }, + }); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(calls).toHaveLength(2); + expect(waits).toEqual([2000]); + }); + + it('clamps an absurd retry-after', async () => { + const waits: number[] = []; + const { transport } = stubTransport([ + { ok: false, status: 503, body: 'down', headers: { 'retry-after': '3600' } }, + { ok: true, status: 200, body: ANSWER }, + ]); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + fetch: transport, + timeoutMs: 30_000, + sleep: async (ms) => { + waits.push(ms); + }, + }); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(waits).toEqual([10_000]); + }); + + it('gives up instead of waiting past its timeout for a retry', async () => { + const waits: number[] = []; + const { transport, calls } = stubTransport([ + { ok: false, status: 429, body: 'slow down', headers: { 'retry-after': '8' } }, + { ok: true, status: 200, body: ANSWER }, + ]); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + fetch: transport, + timeoutMs: 4_000, + sleep: async (ms) => { + waits.push(ms); + }, + }); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'rate_limited', + }); + expect(calls).toHaveLength(1); + expect(waits).toEqual([]); + }); + + it('stops backing off as soon as the caller aborts', async () => { + const controller = new AbortController(); + const { transport, calls } = stubTransport([ + { ok: false, status: 503, body: 'down', headers: { 'retry-after': '2' } }, + { ok: true, status: 200, body: ANSWER }, + ]); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + fetch: transport, + timeoutMs: 30_000, + }); + const started = Date.now(); + setTimeout(() => controller.abort(), 20); + + await expect( + classifier.classify({ state: 'x', questions: QUESTION, signal: controller.signal }), + ).rejects.toMatchObject({ failure: 'aborted' }); + expect(Date.now() - started).toBeLessThan(1_000); + expect(calls).toHaveLength(1); + }); + + it('gives up after maxRetries and names the provider', async () => { + const { classifier, calls } = build([{ ok: false, status: 500, body: 'boom' }], { + maxRetries: 2, + }); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + name: 'ClassificationError', + failure: 'server_error', + provider: 'http', + status: 500, + }); + expect(calls).toHaveLength(3); + }); + + it('does not retry a rejected request', async () => { + const { classifier, calls } = build([{ ok: false, status: 422, body: 'bad question' }]); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'bad_request', + }); + expect(calls).toHaveLength(1); + }); + + it('does not retry a rejected key', async () => { + const { classifier, calls } = build([{ ok: false, status: 401, body: 'nope' }]); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'unauthorized', + }); + expect(calls).toHaveLength(1); + }); + + it('reports a body that is not JSON as malformed', async () => { + const { classifier } = build([{ ok: true, status: 200, body: '504' }]); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'malformed_response', + }); + }); + + it('reports a JSON body with no answers as malformed', async () => { + const { classifier } = build([ + { ok: true, status: 200, body: JSON.stringify({ model: 'test-1' }) }, + ]); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'malformed_response', + }); + }); + + it('drops an answer it cannot read rather than inventing one', async () => { + const body = JSON.stringify({ + model: 'test-1', + answers: { verdict: { type: 'something_new', value: 3 } }, + usage: { input_tokens: 1, output_tokens: 1 }, + }); + const { classifier } = build([{ ok: true, status: 200, body }]); + + const result = await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(result.answers.verdict).toBeUndefined(); + }); + + it('defaults usage when the response omits it', async () => { + const { classifier } = build([ + { + ok: true, + status: 200, + body: JSON.stringify({ + model: 'test-1', + answers: { verdict: { type: 'boolean', probability: 1 } }, + }), + }, + ]); + + const result = await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(result.usage).toEqual({ inputTokens: 0, outputTokens: 0 }); + }); + + it('times out a transport that never settles', async () => { + const transport: ProviderFetch = (_url, init) => + new Promise((_resolve, reject) => { + init.signal.addEventListener('abort', () => reject(new Error('aborted'))); + }); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + timeoutMs: 15, + maxRetries: 0, + fetch: transport, + }); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'timeout', + }); + }); + + it('reports a caller abort as aborted, not as a timeout', async () => { + const controller = new AbortController(); + const transport: ProviderFetch = (_url, init) => + new Promise((_resolve, reject) => { + init.signal.addEventListener('abort', () => reject(new Error('aborted'))); + }); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + timeoutMs: 5_000, + maxRetries: 0, + fetch: transport, + }); + + const pending = classifier.classify({ + state: 'x', + questions: QUESTION, + signal: controller.signal, + }); + controller.abort(); + + await expect(pending).rejects.toMatchObject({ failure: 'aborted' }); + }); + + it('retries a transport-level network failure', async () => { + const { classifier, calls } = build([ + new Error('ECONNRESET'), + { ok: true, status: 200, body: ANSWER }, + ]); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(calls).toHaveLength(2); + }); + + it('refuses to build without a key', () => { + expect(() => createHttpClassifier({ apiKey: ' ', endpoint: ENDPOINT })).toThrow( + ClassificationError, + ); + }); + + it('refuses to build without an endpoint', () => { + expect(() => createHttpClassifier({ apiKey: 'sk-test', endpoint: '' })).toThrow( + ClassificationError, + ); + }); +}); diff --git a/packages/api/src/classification/providers/http.ts b/packages/api/src/classification/providers/http.ts new file mode 100644 index 00000000000..465dbe16bc0 --- /dev/null +++ b/packages/api/src/classification/providers/http.ts @@ -0,0 +1,114 @@ +import type { + Classifier, + ClassificationUsage, + ClassificationAnswer, + ClassificationResult, + ClassificationRequest, +} from '../types'; +import type { ProviderFetch } from './transport'; +import type { Dialect } from './dialect'; +import { toWireQuestion, readAnswer } from './dialect'; +import { ClassificationError } from '../types'; +import { createTransport } from './transport'; + +export const PROVIDER_ID = 'http'; + +export interface HttpProviderOptions { + providerId?: string; + apiKey: string; + /** Full URL, not a base path. */ + endpoint: string; + model?: string; + /** Which wire vocabulary the endpoint speaks. */ + dialect?: Dialect; + /** Nests `state` and `questions` under this key, for hosts that wrap them. */ + requestKey?: string; + /** Reads the answer envelope from this key, for hosts that wrap the response. */ + responseKey?: string; + timeoutMs?: number; + maxRetries?: number; + fetch?: ProviderFetch; + sleep?: (ms: number) => Promise; +} + +export function parseEnvelope( + body: string, + providerId: string, + readOne: (answer: unknown) => ClassificationAnswer | null, + responseKey?: string, +): ClassificationResult { + let parsed: unknown; + try { + parsed = JSON.parse(body); + } catch { + throw new ClassificationError('malformed_response', 'response was not JSON', { + provider: providerId, + }); + } + if (parsed == null || typeof parsed !== 'object') { + throw new ClassificationError('malformed_response', 'response was not an object', { + provider: providerId, + }); + } + const unwrapped = + responseKey != null && responseKey !== '' + ? ((parsed as Record)[responseKey] ?? parsed) + : parsed; + const record = unwrapped as { model?: unknown; answers?: unknown; usage?: unknown }; + if (record.answers == null || typeof record.answers !== 'object') { + throw new ClassificationError('malformed_response', 'response carried no answers', { + provider: providerId, + }); + } + + const answers: Record = {}; + for (const [id, answer] of Object.entries(record.answers as Record)) { + const mapped = readOne(answer); + if (mapped != null) { + answers[id] = mapped; + } + } + + const raw = (record.usage ?? {}) as { input_tokens?: number; output_tokens?: number }; + const usage: ClassificationUsage = { + inputTokens: raw.input_tokens ?? 0, + outputTokens: raw.output_tokens ?? 0, + }; + + return { + model: typeof record.model === 'string' ? record.model : 'unknown', + answers, + usage, + }; +} + +export function createHttpClassifier(options: HttpProviderOptions): Classifier { + const providerId = options.providerId ?? PROVIDER_ID; + const send = createTransport({ ...options, providerId }); + const model = options.model ?? ''; + const dialect: Dialect = options.dialect ?? 'port'; + const { requestKey, responseKey } = options; + + return { + id: providerId, + model, + async classify(request: ClassificationRequest): Promise { + const questions: Record = {}; + for (const [id, question] of Object.entries(request.questions)) { + questions[id] = toWireQuestion(question, dialect); + } + const inner = { state: request.state, questions }; + const payload = JSON.stringify({ + ...(model ? { model } : {}), + ...(requestKey != null && requestKey !== '' ? { [requestKey]: inner } : inner), + }); + const body = await send( + payload, + request.signal, + request.label ?? 'classify', + request.timeoutMs, + ); + return parseEnvelope(body, providerId, (a) => readAnswer(a, dialect), responseKey); + }, + }; +} diff --git a/packages/api/src/classification/providers/transport.ts b/packages/api/src/classification/providers/transport.ts new file mode 100644 index 00000000000..476b96b2b01 --- /dev/null +++ b/packages/api/src/classification/providers/transport.ts @@ -0,0 +1,214 @@ +import { logger } from '@librechat/data-schemas'; +import { ClassificationError } from '../types'; + +const DEFAULT_TIMEOUT_MS = 4_000; +const DEFAULT_MAX_RETRIES = 2; +const BACKOFF_MS = [250, 750, 1_500, 3_000, 6_000]; +const MAX_RETRY_AFTER_MS = 10_000; + +export type ProviderFetch = ( + input: string, + init: { + method: string; + headers: Record; + body: string; + signal: AbortSignal; + }, +) => Promise<{ + ok: boolean; + status: number; + headers: { get(name: string): string | null }; + text(): Promise; +}>; + +export interface TransportOptions { + providerId: string; + apiKey: string; + /** Full URL, not a base path. */ + endpoint: string; + timeoutMs?: number; + maxRetries?: number; + fetch?: ProviderFetch; + sleep?: (ms: number, signal?: AbortSignal) => Promise; +} + +export type Transport = ( + payload: string, + signal: AbortSignal | undefined, + label: string, + timeoutOverrideMs?: number, +) => Promise; + +function failureForStatus(status: number) { + if (status === 401 || status === 403) { + return 'unauthorized' as const; + } + if (status === 429) { + return 'rate_limited' as const; + } + if (status >= 500) { + return 'server_error' as const; + } + return 'bad_request' as const; +} + +function isRetryable(failure: string): boolean { + return failure === 'rate_limited' || failure === 'server_error' || failure === 'network'; +} + +function retryAfterMs(header: string | null): number | undefined { + if (!header) { + return undefined; + } + const seconds = Number(header); + if (Number.isFinite(seconds) && seconds >= 0) { + return Math.min(seconds * 1_000, MAX_RETRY_AFTER_MS); + } + const at = Date.parse(header); + if (Number.isNaN(at)) { + return undefined; + } + return Math.min(Math.max(at - Date.now(), 0), MAX_RETRY_AFTER_MS); +} + +function briefly(body: string): string { + const flat = body.replace(/\s+/g, ' ').trim(); + return flat.length > 200 ? `${flat.slice(0, 200)}…` : flat; +} + +function defaultSleep(ms: number, signal?: AbortSignal): Promise { + if (signal?.aborted === true) { + return Promise.resolve(); + } + return new Promise((resolve) => { + const done = () => { + clearTimeout(timer); + signal?.removeEventListener('abort', done); + resolve(); + }; + const timer = setTimeout(done, ms); + signal?.addEventListener('abort', done, { once: true }); + }); +} + +export function createTransport(options: TransportOptions): Transport { + const { providerId } = options; + const apiKey = options.apiKey?.trim(); + if (!apiKey) { + throw new ClassificationError('unauthorized', 'classifier requires an API key', { + provider: providerId, + }); + } + + const endpoint = options.endpoint?.trim(); + if (!endpoint) { + throw new ClassificationError('bad_request', 'classifier requires a baseURL', { + provider: providerId, + }); + } + + const defaultTimeoutMs = options.timeoutMs ?? DEFAULT_TIMEOUT_MS; + const maxRetries = options.maxRetries ?? DEFAULT_MAX_RETRIES; + const sleep = options.sleep ?? defaultSleep; + const candidate = options.fetch ?? (globalThis.fetch as unknown as ProviderFetch | undefined); + if (typeof candidate !== 'function') { + throw new ClassificationError('network', 'no fetch implementation available', { + provider: providerId, + }); + } + const fetchImpl: ProviderFetch = candidate; + + async function attempt( + payload: string, + signal: AbortSignal | undefined, + timeoutMs: number, + ): Promise { + const timeout = new AbortController(); + const timer = setTimeout(() => timeout.abort(), timeoutMs); + const combined = signal != null ? AbortSignal.any([signal, timeout.signal]) : timeout.signal; + + try { + const response = await fetchImpl(endpoint, { + method: 'POST', + headers: { + Authorization: `Bearer ${apiKey}`, + 'Content-Type': 'application/json', + }, + body: payload, + signal: combined, + }); + + const text = await response.text(); + if (!response.ok) { + const error = new ClassificationError( + failureForStatus(response.status), + `classifier returned ${response.status}: ${briefly(text)}`, + { provider: providerId, status: response.status }, + ); + const wait = retryAfterMs(response.headers.get('retry-after')); + if (wait != null) { + error.retryAfterMs = wait; + } + throw error; + } + return text; + } catch (error) { + if (error instanceof ClassificationError) { + throw error; + } + if (signal?.aborted === true) { + throw new ClassificationError('aborted', 'caller aborted the request', { + provider: providerId, + }); + } + if (timeout.signal.aborted) { + throw new ClassificationError('timeout', `no answer within ${timeoutMs}ms`, { + provider: providerId, + }); + } + const message = error instanceof Error ? error.message : String(error); + throw new ClassificationError('network', message, { provider: providerId }); + } finally { + clearTimeout(timer); + } + } + + /** `timeoutMs` bounds the whole call, retries and backoff included, not each attempt. */ + return async function send(payload, signal, label, timeoutOverrideMs) { + const timeoutMs = + timeoutOverrideMs != null && timeoutOverrideMs > 0 ? timeoutOverrideMs : defaultTimeoutMs; + const deadline = Date.now() + timeoutMs; + let lastError: ClassificationError | undefined; + for (let attemptNo = 0; attemptNo <= maxRetries; attemptNo++) { + try { + const started = Date.now(); + const body = await attempt(payload, signal, Math.max(1, deadline - started)); + logger.debug(`[classification] ${label} answered in ${Date.now() - started}ms`); + return body; + } catch (error) { + lastError = + error instanceof ClassificationError + ? error + : new ClassificationError('network', String(error), { provider: providerId }); + if (attemptNo === maxRetries || !isRetryable(lastError.failure)) { + break; + } + const wait = + lastError.retryAfterMs ?? BACKOFF_MS[Math.min(attemptNo, BACKOFF_MS.length - 1)]; + if (wait >= deadline - Date.now()) { + break; + } + await sleep(wait, signal); + if (signal?.aborted === true) { + lastError = new ClassificationError('aborted', 'caller aborted the request', { + provider: providerId, + }); + break; + } + } + } + throw ( + lastError ?? new ClassificationError('network', 'request failed', { provider: providerId }) + ); + }; +} diff --git a/packages/api/src/classification/questions.spec.ts b/packages/api/src/classification/questions.spec.ts new file mode 100644 index 00000000000..306c3611d7d --- /dev/null +++ b/packages/api/src/classification/questions.spec.ts @@ -0,0 +1,61 @@ +import type { ChoiceAnswer, BooleanAnswer } from './types'; +import { choice, isTrue, ranked, boolean } from './questions'; + +describe('question builders', () => { + it('builds a boolean without criteria', () => { + expect(boolean('Is this urgent?')).toEqual({ + type: 'boolean', + instructions: 'Is this urgent?', + }); + }); + + it('builds a boolean with criteria', () => { + expect(boolean('Is this urgent?', { true: 'yes means', false: 'no means' })).toEqual({ + type: 'boolean', + instructions: 'Is this urgent?', + criteria: { true: 'yes means', false: 'no means' }, + }); + }); + + it('builds a choice', () => { + expect(choice('Which team?', { billing: 'money', tech: 'bugs' })).toEqual({ + type: 'choice', + instructions: 'Which team?', + criteria: { billing: 'money', tech: 'bugs' }, + }); + }); +}); + +describe('isTrue', () => { + const answer: BooleanAnswer = { type: 'boolean', probability: 0.8 }; + + it('compares against the threshold inclusively', () => { + expect(isTrue(answer, 0.8)).toBe(true); + expect(isTrue(answer, 0.81)).toBe(false); + }); + + it('is false for a missing answer', () => { + expect(isTrue(undefined, 0)).toBe(false); + }); +}); + +describe('choice helpers', () => { + const answer: ChoiceAnswer = { + type: 'choice', + choice: 'tech', + confidence: 0.7, + probabilities: { tech: 0.7, billing: 0.2, sales: 0.05 }, + }; + + it('ranks options most probable first', () => { + expect(ranked(answer)).toEqual(['tech', 'billing', 'sales']); + }); + + it('drops options under the floor', () => { + expect(ranked(answer, 0.1)).toEqual(['tech', 'billing']); + }); + + it('handles a missing answer', () => { + expect(ranked(undefined)).toEqual([]); + }); +}); diff --git a/packages/api/src/classification/questions.ts b/packages/api/src/classification/questions.ts new file mode 100644 index 00000000000..86e51bf842c --- /dev/null +++ b/packages/api/src/classification/questions.ts @@ -0,0 +1,38 @@ +import type { + ChoiceAnswer, + BooleanAnswer, + ChoiceQuestion, + BooleanQuestion, + ClassificationText, +} from './types'; + +export function boolean( + instructions: ClassificationText, + criteria?: { true?: ClassificationText; false?: ClassificationText }, +): BooleanQuestion { + return criteria == null + ? { type: 'boolean', instructions } + : { type: 'boolean', instructions, criteria }; +} + +export function choice( + instructions: ClassificationText, + criteria: Record, +): ChoiceQuestion { + return { type: 'choice', instructions, criteria }; +} + +export function isTrue(answer: BooleanAnswer | undefined, threshold: number): boolean { + return answer != null && answer.probability >= threshold; +} + +/** Options above `floor`, most probable first. */ +export function ranked(answer: ChoiceAnswer | undefined, floor = 0): string[] { + if (answer == null) { + return []; + } + return Object.entries(answer.probabilities) + .filter(([, probability]) => probability >= floor) + .sort((a, b) => b[1] - a[1]) + .map(([option]) => option); +} diff --git a/packages/api/src/classification/registry.ts b/packages/api/src/classification/registry.ts new file mode 100644 index 00000000000..16ba40de231 --- /dev/null +++ b/packages/api/src/classification/registry.ts @@ -0,0 +1,80 @@ +import type { TClassificationProviderConfig } from 'librechat-data-provider'; +import type { ProviderFetch } from './providers/transport'; +import type { Classifier } from './types'; +import { createHttpClassifier } from './providers/http'; + +export type ProviderSettings = TClassificationProviderConfig; + +export const DEFAULT_API_KEY_ENV = 'CLASSIFIER_API_KEY'; + +/** + * Known hosts, as settings rather than code. They all serve the same question + * shapes over HTTP and differ only in URL, model name and how the body is + * wrapped, so a new one is an entry here or, for an operator who cannot wait + * for a release, the same fields written in `librechat.yaml`. + */ +export const PRESETS: Record = { + http: { + dialect: 'port', + apiKeyEnv: DEFAULT_API_KEY_ENV, + }, + typesafe: { + baseURL: 'https://api.typesafe.ai/v1/systemone', + model: 'jev-latest', + dialect: 'systemone', + apiKeyEnv: 'TYPESAFE_API_KEY', + }, + openrouter: { + baseURL: 'https://openrouter.ai/api/alpha/decisions', + model: '~typesafe/jev-latest', + dialect: 'systemone', + apiKeyEnv: 'OPENROUTER_KEY', + }, + cloudflare: { + /** No default URL: the account id is part of it. */ + model: 'typesafe/jev', + dialect: 'systemone', + requestKey: 'input', + responseKey: 'result', + apiKeyEnv: 'CLOUDFLARE_API_TOKEN', + }, +}; + +export function presetFor(name: string | undefined): ProviderSettings | null { + if (!name) { + return null; + } + return PRESETS[name] ?? null; +} + +/** Operator settings win over the preset, field by field. */ +export function mergeSettings( + preset: ProviderSettings | null, + configured: ProviderSettings | undefined, +): ProviderSettings { + return { ...(preset ?? {}), ...(configured ?? {}) }; +} + +export function providerNames(): string[] { + return Object.keys(PRESETS); +} + +export function createClassifier( + settings: ProviderSettings, + apiKey: string, + fetch?: ProviderFetch, + providerId?: string, +): Classifier { + return createHttpClassifier({ + providerId, + apiKey, + endpoint: settings.baseURL ?? '', + model: settings.model, + dialect: settings.dialect, + requestKey: settings.requestKey, + responseKey: settings.responseKey, + timeoutMs: settings.timeoutMs, + maxRetries: settings.maxRetries, + fetch, + }); +} diff --git a/packages/api/src/classification/resolve.ts b/packages/api/src/classification/resolve.ts new file mode 100644 index 00000000000..d3a04cc4961 --- /dev/null +++ b/packages/api/src/classification/resolve.ts @@ -0,0 +1,110 @@ +import { logger } from '@librechat/data-schemas'; +import type { TClassificationConfig } from 'librechat-data-provider'; +import type { ProviderFetch } from './providers/transport'; +import type { Classifier } from './types'; +import { + presetFor, + mergeSettings, + createClassifier, + providerNames, + DEFAULT_API_KEY_ENV, +} from './registry'; +import { ClassificationError } from './types'; + +interface CacheEntry { + apiKey: string; + provider: string; + classifier: Classifier; +} + +const cache = new WeakMap(); +const warned = new Set(); + +export interface ResolveClassifierParams { + config?: TClassificationConfig | null; + apiKey?: string; + fetch?: ProviderFetch; +} + +export function resolveClassifier(params: ResolveClassifierParams): Classifier | null { + const config = params.config; + if (config == null || config.enabled !== true) { + return null; + } + + const providerId = config.provider; + const configured = config.providers?.[providerId]; + const preset = presetFor(providerId); + /** An unknown name is fine when the operator described the host in full. */ + if (preset == null && configured?.baseURL == null) { + if (!warned.has(`unknown:${providerId}`)) { + warned.add(`unknown:${providerId}`); + logger.warn( + `[classification] provider "${providerId}" has no preset and no baseURL. ` + + `Known presets: ${providerNames().join(', ')}. Classification stays off.`, + ); + } + return null; + } + + const settings = mergeSettings(preset, configured); + if (!settings.baseURL) { + if (!warned.has(`url:${providerId}`)) { + warned.add(`url:${providerId}`); + logger.warn( + `[classification] provider "${providerId}" needs a baseURL; classification stays off.`, + ); + } + return null; + } + const apiKeyEnv = settings.apiKeyEnv ?? DEFAULT_API_KEY_ENV; + const apiKey = (params.apiKey ?? process.env[apiKeyEnv] ?? '').trim(); + if (!apiKey) { + if (!warned.has(`key:${providerId}`)) { + warned.add(`key:${providerId}`); + logger.warn( + `[classification] provider "${providerId}" is configured but ${apiKeyEnv} is not set; ` + + 'every classification capability stays off.', + ); + } + return null; + } + + const cached = cache.get(config); + if (cached != null && cached.apiKey === apiKey && cached.provider === providerId) { + return cached.classifier; + } + + try { + const classifier = createClassifier(settings, apiKey, params.fetch, providerId); + cache.set(config, { apiKey, provider: providerId, classifier }); + return classifier; + } catch (error) { + logger.error( + `[classification] could not build provider "${providerId}"; capabilities stay off`, + error instanceof ClassificationError ? { failure: error.failure } : error, + ); + return null; + } +} + +export type ClassificationCapabilityName = 'toolSelection' | 'memoryGate'; + +export function classificationCapability( + config: TClassificationConfig | null | undefined, + capability: K, + options?: { apiKey?: string; fetch?: ProviderFetch }, +): { classifier: Classifier; settings: NonNullable[K] } | null { + if (config == null || config.enabled !== true) { + return null; + } + const settings = config[capability]; + if (settings == null || settings.enabled !== true) { + return null; + } + const classifier = resolveClassifier({ config, apiKey: options?.apiKey, fetch: options?.fetch }); + if (classifier == null) { + return null; + } + return { classifier, settings }; +} diff --git a/packages/api/src/classification/types.ts b/packages/api/src/classification/types.ts new file mode 100644 index 00000000000..a616ca60de2 --- /dev/null +++ b/packages/api/src/classification/types.ts @@ -0,0 +1,110 @@ +export type ClassificationJson = + | string + | number + | boolean + | null + | ClassificationJson[] + | { [key: string]: ClassificationJson }; + +export type ClassificationText = + | string + | ClassificationJson[] + | { [key: string]: ClassificationJson }; + +export type ClassificationState = ClassificationText; + +export interface BooleanQuestion { + type: 'boolean'; + instructions: ClassificationText; + criteria?: { + true?: ClassificationText; + false?: ClassificationText; + }; +} + +export interface ChoiceQuestion { + type: 'choice'; + instructions: ClassificationText; + criteria: Record; +} + +export type ClassificationQuestion = BooleanQuestion | ChoiceQuestion; + +export interface BooleanAnswer { + type: 'boolean'; + probability: number; +} + +export interface ChoiceAnswer { + type: 'choice'; + choice: string; + /** `null` when the provider cannot measure it, which is not the same as 0. */ + confidence: number | null; + probabilities: Record; +} + +export type ClassificationAnswer = BooleanAnswer | ChoiceAnswer; + +export interface ClassificationUsage { + inputTokens: number; + outputTokens: number; +} + +export interface ClassificationRequest { + state: ClassificationState; + questions: Record; + signal?: AbortSignal; + label?: string; + /** Overrides the provider's timeout for this request alone. */ + timeoutMs?: number; +} + +export interface ClassificationResult { + model: string; + answers: Record; + usage: ClassificationUsage; +} + +export interface Classifier { + readonly id: string; + readonly model: string; + classify(request: ClassificationRequest): Promise; +} + +export type ClassificationFailure = + | 'timeout' + | 'aborted' + | 'rate_limited' + | 'unauthorized' + | 'bad_request' + | 'server_error' + | 'network' + | 'unsupported_question' + | 'malformed_response'; + +export class ClassificationError extends Error { + readonly failure: ClassificationFailure; + readonly provider: string; + readonly status?: number; + retryAfterMs?: number; + + constructor( + failure: ClassificationFailure, + message: string, + options?: { provider?: string; status?: number }, + ) { + super(message); + this.name = 'ClassificationError'; + this.failure = failure; + this.provider = options?.provider ?? 'unknown'; + this.status = options?.status; + } +} + +export function isBooleanAnswer(answer: ClassificationAnswer | undefined): answer is BooleanAnswer { + return answer?.type === 'boolean'; +} + +export function isChoiceAnswer(answer: ClassificationAnswer | undefined): answer is ChoiceAnswer { + return answer?.type === 'choice'; +} diff --git a/packages/api/src/index.ts b/packages/api/src/index.ts index 83a0cc9b63d..0113e5a5e44 100644 --- a/packages/api/src/index.ts +++ b/packages/api/src/index.ts @@ -97,6 +97,8 @@ export * from './images'; export * from './storage'; /* Tools */ export * from './tools'; +/* Classification */ +export * from './classification'; /* web search */ export * from './web'; /* Langfuse */ diff --git a/packages/api/src/memory/gate.spec.ts b/packages/api/src/memory/gate.spec.ts new file mode 100644 index 00000000000..6d3de34fd75 --- /dev/null +++ b/packages/api/src/memory/gate.spec.ts @@ -0,0 +1,370 @@ +import { HumanMessage, AIMessage } from '@librechat/agents/langchain/messages'; +import type { + Classifier, + ClassificationResult, + ClassificationRequest, +} from '~/classification/types'; +import type { MemoryGateSettings } from './gate'; +import { createMemoryGate, transcribeTail, MEMORY_REQUEST_QUESTION } from './gate'; + +const ON: MemoryGateSettings = { + enabled: true, + threshold: 0.25, + categorize: false, + categoryThreshold: 0.4, + detectUpdates: false, +}; + +function stubClassifier(probability: number | Error): { + classifier: Classifier; + requests: ClassificationRequest[]; +} { + const requests: ClassificationRequest[] = []; + const classifier: Classifier = { + id: 'stub', + model: 'stub-1', + async classify(params: ClassificationRequest): Promise { + requests.push(params); + if (probability instanceof Error) { + throw probability; + } + return { + model: 'stub-1', + answers: { request: { type: 'boolean', probability } }, + usage: { inputTokens: 40, outputTokens: 4 }, + }; + }, + }; + return { classifier, requests }; +} + +const TURN = [ + new HumanMessage('I always deploy on Fridays, never on Mondays.'), + new AIMessage('Noted.'), +]; + +describe('transcribeTail', () => { + it('reads the last messages in order, labelled by role', () => { + const text = transcribeTail(TURN, 6, 1_000); + + expect(text).toBe('human: I always deploy on Fridays, never on Mondays.\nai: Noted.'); + }); + + it('keeps only the newest messages inside the window', () => { + const messages = [ + new HumanMessage('oldest'), + new HumanMessage('middle'), + new HumanMessage('newest'), + ]; + + expect(transcribeTail(messages, 2, 1_000)).toBe('human: middle\nhuman: newest'); + }); + + it('stays inside the character budget', () => { + const messages = [new HumanMessage('x'.repeat(500)), new HumanMessage('y'.repeat(500))]; + + expect(transcribeTail(messages, 10, 100).length).toBeLessThanOrEqual(100); + }); + + it('reads text out of structured content parts', () => { + const message = new HumanMessage({ + content: [ + { type: 'text', text: 'I use metric units' }, + { type: 'image_url', image_url: { url: 'https://example.test/a.png' } }, + ], + }); + + expect(transcribeTail([message], 4, 1_000)).toBe('human: I use metric units'); + }); + + it('skips messages with no text at all', () => { + const empty = new HumanMessage({ content: [] }); + + expect(transcribeTail([empty, new HumanMessage('real')], 4, 1_000)).toBe('human: real'); + }); +}); + +describe('createMemoryGate', () => { + it('is absent when the capability is off, so the caller keeps one path', () => { + const { classifier } = stubClassifier(0.9); + + expect(createMemoryGate({ classifier, settings: { ...ON, enabled: false } })).toBeNull(); + }); + + it('processes a turn that asks to remember something', async () => { + const { classifier, requests } = stubClassifier(0.88); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.({ messages: TURN })).resolves.toMatchObject({ process: true }); + expect(requests).toHaveLength(1); + expect(requests[0].questions.request).toBeDefined(); + }); + + it('skips a turn that asks for no memory change', async () => { + const { classifier } = stubClassifier(0.03); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.({ messages: [new HumanMessage('thanks!')] })).resolves.toMatchObject({ + process: false, + }); + }); + + it('treats the threshold as inclusive', async () => { + const { classifier } = stubClassifier(0.25); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.({ messages: TURN })).resolves.toMatchObject({ process: true }); + }); + + it('respects a stricter threshold', async () => { + const { classifier } = stubClassifier(0.5); + const gate = createMemoryGate({ classifier, settings: { ...ON, threshold: 0.8 } }); + + await expect(gate?.({ messages: TURN })).resolves.toMatchObject({ process: false }); + }); + + it('processes the turn when the judgment fails, rather than losing a memory', async () => { + const { classifier } = stubClassifier(new Error('upstream exploded')); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.({ messages: TURN })).resolves.toMatchObject({ process: true }); + }); + + it('processes the turn when the answer comes back the wrong shape', async () => { + const classifier: Classifier = { + id: 'stub', + model: 'stub-1', + classify: async () => ({ + model: 'stub-1', + answers: { request: { type: 'choice', choice: 'yes', confidence: 1, probabilities: {} } }, + usage: { inputTokens: 1, outputTokens: 1 }, + }), + }; + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.({ messages: TURN })).resolves.toMatchObject({ process: true }); + }); + + it('skips an empty turn without calling out', async () => { + const { classifier, requests } = stubClassifier(0.9); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.({ messages: [] })).resolves.toMatchObject({ process: false }); + expect(requests).toHaveLength(0); + }); + + it('judges the newest user message and sends earlier turns only as context', async () => { + const { classifier, requests } = stubClassifier(0.9); + const gate = createMemoryGate({ classifier, settings: ON }); + + await gate?.({ + messages: [ + new HumanMessage('I prefer answers in Japanese from now on'), + new AIMessage('Understood.'), + new HumanMessage('take a screenshot of the page'), + ], + }); + + expect(requests[0].state).toEqual({ + latest: 'take a screenshot of the page', + conversation: 'human: I prefer answers in Japanese from now on\nai: Understood.', + }); + }); + + it('skips without calling out when the window holds no user message', async () => { + const { classifier, requests } = stubClassifier(0.9); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.({ messages: [new AIMessage('Hello!')] })).resolves.toMatchObject({ + process: false, + }); + expect(requests).toHaveLength(0); + }); +}); + +describe('createMemoryGate usage', () => { + it('reports what each judgment cost', async () => { + const { classifier } = stubClassifier(0.9); + const onUsage = jest.fn(); + const gate = createMemoryGate({ classifier, settings: ON, onUsage }); + + await gate?.({ messages: TURN }); + + expect(onUsage).toHaveBeenCalledWith({ inputTokens: 40, outputTokens: 4 }, 'stub-1'); + }); + + it('reports nothing when the judgment fails', async () => { + const { classifier } = stubClassifier(new Error('upstream exploded')); + const onUsage = jest.fn(); + const gate = createMemoryGate({ classifier, settings: ON, onUsage }); + + await gate?.({ messages: TURN }); + + expect(onUsage).not.toHaveBeenCalled(); + }); +}); + +describe('createMemoryGate prompt overrides', () => { + it('uses the built-in wording when nothing is configured', async () => { + const { classifier, requests } = stubClassifier(0.9); + const gate = createMemoryGate({ classifier, settings: ON }); + + await gate?.({ messages: TURN }); + + const question = requests[0].questions.request as unknown as { + instructions: string; + criteria: { true: string; false: string }; + }; + expect(question.instructions).toBe(MEMORY_REQUEST_QUESTION.instructions); + expect(question.criteria.true).toBe(MEMORY_REQUEST_QUESTION.criteria?.true); + expect(question.criteria.false).toBe(MEMORY_REQUEST_QUESTION.criteria?.false); + }); + + it('sends operator wording instead when it is configured', async () => { + const { classifier, requests } = stubClassifier(0.9); + const gate = createMemoryGate({ + classifier, + settings: { + ...ON, + instructions: 'Is there a dietary requirement here?', + whenTrue: 'An allergy or a standing preference.', + whenFalse: 'Anything about one meal only.', + }, + }); + + await gate?.({ messages: TURN }); + + const question = requests[0].questions.request as unknown as { + instructions: string; + criteria: { true: string; false: string }; + }; + expect(question.instructions).toBe('Is there a dietary requirement here?'); + expect(question.criteria.true).toBe('An allergy or a standing preference.'); + expect(question.criteria.false).toBe('Anything about one meal only.'); + }); +}); + +describe('memory classification', () => { + const KEYS = ['work_context', 'preferences', 'personal']; + + function stubAnswers(answers: ClassificationResult['answers']): { + classifier: Classifier; + requests: ClassificationRequest[]; + } { + const requests: ClassificationRequest[] = []; + const classifier: Classifier = { + id: 'stub', + model: 'stub-1', + async classify(params: ClassificationRequest): Promise { + requests.push(params); + return { model: 'stub-1', answers, usage: { inputTokens: 1, outputTokens: 1 } }; + }, + }; + return { classifier, requests }; + } + + const request = { type: 'boolean' as const, probability: 0.9 }; + + it('asks only the durability question by default', async () => { + const { classifier, requests } = stubAnswers({ request }); + const gate = createMemoryGate({ classifier, settings: ON }); + + await gate?.({ messages: TURN, validKeys: KEYS }); + + expect(Object.keys(requests[0].questions)).toEqual(['request']); + }); + + it('adds the category and update questions to the same request', async () => { + const { classifier, requests } = stubAnswers({ request }); + const gate = createMemoryGate({ + classifier, + settings: { ...ON, categorize: true, detectUpdates: true }, + }); + + await gate?.({ messages: TURN, validKeys: KEYS }); + + expect(requests).toHaveLength(1); + expect(Object.keys(requests[0].questions).sort()).toEqual(['category', 'request', 'updates']); + }); + + it('skips categorization when no valid keys are configured', async () => { + const { classifier, requests } = stubAnswers({ request }); + const gate = createMemoryGate({ classifier, settings: { ...ON, categorize: true } }); + + await gate?.({ messages: TURN }); + + expect(Object.keys(requests[0].questions)).toEqual(['request']); + }); + + it('suggests the winning key', async () => { + const { classifier } = stubAnswers({ + request, + category: { + type: 'choice', + choice: 'work_context', + confidence: 0.8, + probabilities: { work_context: 0.8, preferences: 0.15, personal: 0.05 }, + }, + }); + const gate = createMemoryGate({ classifier, settings: { ...ON, categorize: true } }); + + const judgment = await gate?.({ messages: TURN, validKeys: KEYS }); + + expect(judgment?.hint).toContain('work_context'); + expect(judgment?.hint).toContain('Ignore this'); + }); + + it('drops a key it is not confident about', async () => { + const { classifier } = stubAnswers({ + request, + category: { + type: 'choice', + choice: 'work_context', + confidence: 0.3, + probabilities: { work_context: 0.35, preferences: 0.33, personal: 0.32 }, + }, + }); + const gate = createMemoryGate({ classifier, settings: { ...ON, categorize: true } }); + + const judgment = await gate?.({ messages: TURN, validKeys: KEYS }); + + expect(judgment?.process).toBe(true); + expect(judgment?.hint).toBeUndefined(); + }); + + it('says so when the turn changes an existing fact', async () => { + const { classifier } = stubAnswers({ + request, + category: { + type: 'choice', + choice: 'preferences', + confidence: 0.9, + probabilities: { preferences: 0.9, work_context: 0.05, personal: 0.05 }, + }, + updates: { type: 'boolean', probability: 0.8 }, + }); + const gate = createMemoryGate({ + classifier, + settings: { ...ON, categorize: true, detectUpdates: true }, + }); + + const judgment = await gate?.({ messages: TURN, validKeys: KEYS }); + + expect(judgment?.hint).toContain('a change to what is already stored'); + }); + + it('carries no hint when the turn asks for nothing', async () => { + const { classifier } = stubAnswers({ + request: { type: 'boolean', probability: 0.01 }, + category: { + type: 'choice', + choice: 'personal', + confidence: 1, + probabilities: { personal: 1 }, + }, + }); + const gate = createMemoryGate({ classifier, settings: { ...ON, categorize: true } }); + + expect(await gate?.({ messages: TURN, validKeys: KEYS })).toEqual({ process: false }); + }); +}); diff --git a/packages/api/src/memory/gate.ts b/packages/api/src/memory/gate.ts new file mode 100644 index 00000000000..7d4cdf547c6 --- /dev/null +++ b/packages/api/src/memory/gate.ts @@ -0,0 +1,231 @@ +import { logger } from '@librechat/data-schemas'; +import type { BaseMessage } from '@librechat/agents/langchain/messages'; +import type { TClassificationConfig } from 'librechat-data-provider'; +import type { + Classifier, + BooleanQuestion, + ClassificationUsage, + ClassificationQuestion, +} from '~/classification/types'; +import { isBooleanAnswer, isChoiceAnswer } from '~/classification/types'; +import { boolean, choice, isTrue } from '~/classification/questions'; + +export interface MemoryJudgment { + process: boolean; + hint?: string; +} + +export interface MemoryGateInput { + messages: BaseMessage[]; + /** Keys the memory model may write. Categorization is skipped without them. */ + validKeys?: string[]; +} + +export type MemoryGate = (input: MemoryGateInput) => Promise; + +export type MemoryGateSettings = TClassificationConfig['memoryGate']; + +export interface CreateMemoryGateParams { + classifier: Classifier; + settings: MemoryGateSettings; + signal?: AbortSignal; + windowSize?: number; + maxChars?: number; + /** Receives what each judgment cost, so the caller can bill it. */ + onUsage?: (usage: ClassificationUsage, model: string) => void; +} + +const DEFAULT_WINDOW = 6; +const DEFAULT_MAX_CHARS = 8_000; +const PROCESS = { process: true } as const; + +/** Matches the default memory instructions, which store only what the user asks to keep. */ +export const MEMORY_REQUEST_QUESTION: BooleanQuestion = boolean( + "Does `latest`, the user's newest message, ask the assistant to remember, update or forget " + + 'something about the user? `conversation` is the earlier context and is only there to help ' + + 'read `latest`.', + { + true: + 'An explicit request to remember, store, change or delete something, including a short ' + + 'yes to an offer to remember.', + false: + 'Anything else, including preferences or facts the user mentions without asking for ' + + 'them to be kept.', + }, +); + +export const CATEGORY_INSTRUCTIONS = 'Which stored memory does the request in `latest` concern?'; + +export const UPDATE_QUESTION: BooleanQuestion = boolean( + 'Does `latest` change something already known about this user, rather than adding ' + + 'something new?', + { + true: 'It corrects, replaces or narrows a fact the assistant would already hold.', + false: 'It is new, or it repeats what is already stored without changing it.', + }, +); + +function messageText(message: BaseMessage): string { + const content = message.content; + if (typeof content === 'string') { + return content; + } + if (!Array.isArray(content)) { + return ''; + } + const parts: string[] = []; + for (const part of content) { + if (typeof part === 'string') { + parts.push(part); + continue; + } + if (part != null && typeof part === 'object' && 'text' in part) { + const text = (part as { text?: unknown }).text; + if (typeof text === 'string') { + parts.push(text); + } + } + } + return parts.join(' '); +} + +export function transcribeTail( + messages: BaseMessage[], + windowSize: number, + maxChars: number, +): string { + const lines: string[] = []; + let used = 0; + for (let i = messages.length - 1; i >= 0 && lines.length < windowSize; i--) { + const message = messages[i]; + const text = messageText(message).replace(/\s+/g, ' ').trim(); + if (text.length === 0) { + continue; + } + const role = message._getType?.() ?? 'message'; + const line = `${role}: ${text}`; + const clipped = + line.length > maxChars - used ? line.slice(0, Math.max(0, maxChars - used)) : line; + if (clipped.length === 0) { + break; + } + lines.push(clipped); + used += clipped.length; + if (used >= maxChars) { + break; + } + } + return lines.reverse().join('\n'); +} + +/** Index of the newest user message, or -1 when the window holds none. */ +function latestHumanIndex(messages: BaseMessage[]): number { + for (let i = messages.length - 1; i >= 0; i--) { + if (messages[i]._getType?.() === 'human' && messageText(messages[i]).trim().length > 0) { + return i; + } + } + return -1; +} + +export function buildHint(key: string | null, updates: boolean): string | undefined { + if (key == null) { + return undefined; + } + const change = updates + ? ' It looks like a change to what is already stored there, not a new fact.' + : ''; + return ( + '\n' + + `This request most likely belongs under \`${key}\`.${change}\n` + + 'Ignore this if it does not fit what the user actually said.\n' + + '' + ); +} + +export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | null { + const { classifier, settings, signal, onUsage } = params; + if (settings?.enabled !== true) { + return null; + } + const windowSize = params.windowSize ?? DEFAULT_WINDOW; + const maxChars = params.maxChars ?? DEFAULT_MAX_CHARS; + const threshold = settings.threshold; + const request = boolean(settings.instructions ?? MEMORY_REQUEST_QUESTION.instructions, { + true: settings.whenTrue ?? MEMORY_REQUEST_QUESTION.criteria?.true, + false: settings.whenFalse ?? MEMORY_REQUEST_QUESTION.criteria?.false, + }); + + function buildQuestions(validKeys: string[]): Record { + const questions: Record = { request }; + if (settings.categorize === true && validKeys.length > 0) { + questions.category = choice( + CATEGORY_INSTRUCTIONS, + Object.fromEntries(validKeys.map((key) => [key, null])), + ); + } + if (settings.detectUpdates === true) { + questions.updates = UPDATE_QUESTION; + } + return questions; + } + + return async function memoryGate({ + messages, + validKeys = [], + }: MemoryGateInput): Promise { + const window = messages ?? []; + const index = latestHumanIndex(window); + if (index === -1) { + return { process: false }; + } + const latest = messageText(window[index]).replace(/\s+/g, ' ').trim().slice(0, maxChars); + const transcript = transcribeTail( + window.slice(0, index), + windowSize - 1, + Math.max(0, maxChars - latest.length), + ); + try { + const response = await classifier.classify({ + label: 'memory-gate', + signal, + timeoutMs: settings.timeoutMs, + state: { latest, conversation: transcript }, + questions: buildQuestions(validKeys), + }); + onUsage?.(response.usage, classifier.model); + + const answer = response.answers.request; + if (!isBooleanAnswer(answer)) { + return PROCESS; + } + if (!isTrue(answer, threshold)) { + logger.debug( + `[memoryGate] request ${answer.probability.toFixed(2)} below ${threshold}: skipping`, + ); + return { process: false }; + } + + const categoryAnswer = response.answers.category; + const key = + isChoiceAnswer(categoryAnswer) && + (categoryAnswer.probabilities[categoryAnswer.choice] ?? 0) >= settings.categoryThreshold + ? categoryAnswer.choice + : null; + const updatesAnswer = response.answers.updates; + const updates = isBooleanAnswer(updatesAnswer) && updatesAnswer.probability >= 0.5; + + logger.debug( + `[memoryGate] request ${answer.probability.toFixed(2)}: processing` + + (key != null ? `, suggesting \`${key}\`${updates ? ' as an update' : ''}` : ''), + ); + return { process: true, hint: buildHint(key, updates) }; + } catch (error) { + logger.warn( + '[memoryGate] judgment failed, processing memory as usual: ' + + (error instanceof Error ? error.message : String(error)), + ); + return PROCESS; + } + }; +} diff --git a/packages/api/src/memory/index.ts b/packages/api/src/memory/index.ts index c916eded625..5c7cf4b3f3d 100644 --- a/packages/api/src/memory/index.ts +++ b/packages/api/src/memory/index.ts @@ -2,3 +2,4 @@ export * from './config'; export * from './authorization'; export * from './handlers'; export * from './protection'; +export * from './gate'; diff --git a/packages/api/src/tools/index.ts b/packages/api/src/tools/index.ts index d901755b7cd..c9c4e393871 100644 --- a/packages/api/src/tools/index.ts +++ b/packages/api/src/tools/index.ts @@ -6,3 +6,4 @@ export * from './toolkits'; export * from './definitions'; export * from './classification'; export * from './rolePermissions'; +export * from './predict'; diff --git a/packages/api/src/tools/predict.spec.ts b/packages/api/src/tools/predict.spec.ts new file mode 100644 index 00000000000..d174885f7c8 --- /dev/null +++ b/packages/api/src/tools/predict.spec.ts @@ -0,0 +1,531 @@ +import { classificationSchema } from 'librechat-data-provider'; +import { HumanMessage } from '@librechat/agents/langchain/messages'; +import type { LCToolRegistry, LCTool } from '@librechat/agents'; +import type { + Classifier, + ClassificationResult, + ClassificationRequest, +} from '~/classification/types'; +import type { PredictCandidate, ToolSelectionConfig } from './predict'; +import { + predictTools, + shortlistSize, + namedInRequest, + batchCandidates, + deferredCandidates, + predictToolsForTurn, + RANKING_QUESTION, + NEEDS_TOOL_QUESTION, +} from './predict'; + +const CONFIG: ToolSelectionConfig = { + enabled: true, + shortlist: 3, + minProbability: 0.05, + needsToolThreshold: 0.15, + lowConfidenceExtra: 0, + lowConfidenceBelow: 0.5, + surfaceNamedTools: false, + maxCatalogTools: 200, + descriptionChars: 300, +}; + +const NO_MATCH = '__no_tool_fits__'; + +function candidates(...names: string[]): PredictCandidate[] { + return names.map((name) => ({ name, description: `does ${name}` })); +} + +/** Records the requests and answers each with the queued distributions. */ +function stubClient( + distributions: Array>, + needsTool = 0.9, +): { classifier: Classifier; requests: ClassificationRequest[] } { + const requests: ClassificationRequest[] = []; + let index = 0; + const classifier: Classifier = { + id: 'stub', + model: 'stub-1', + async classify(params: ClassificationRequest): Promise { + requests.push(params); + const probabilities = distributions[Math.min(index, distributions.length - 1)]; + index++; + const entries = Object.entries(probabilities).sort((a, b) => b[1] - a[1]); + const answers: ClassificationResult['answers'] = { + best_tool: { + type: 'choice', + choice: entries[0][0], + confidence: entries[0][1], + probabilities, + }, + }; + if (params.questions.needs_tool != null) { + answers.needs_tool = { type: 'boolean', probability: needsTool }; + } + return { + model: 'stub-1', + answers, + usage: { inputTokens: 100, outputTokens: 10 }, + }; + }, + }; + return { classifier, requests }; +} + +describe('batchCandidates', () => { + it('keeps a catalog that fits in one batch', () => { + expect(batchCandidates(candidates('a', 'b', 'c'), 200)).toHaveLength(1); + }); + + it('splits a catalog larger than the batch size', () => { + const many = candidates(...Array.from({ length: 450 }, (_, i) => `tool_${i}`)); + const batches = batchCandidates(many, 200); + expect(batches.map((b) => b.length)).toEqual([200, 200, 50]); + }); + + it('never exceeds the API option ceiling, whatever the config asks for', () => { + const many = candidates(...Array.from({ length: 600 }, (_, i) => `tool_${i}`)); + const batches = batchCandidates(many, 10_000); + for (const batch of batches) { + expect(batch.length).toBeLessThanOrEqual(254); + } + expect(batches.reduce((n, b) => n + b.length, 0)).toBe(600); + }); +}); + +describe('predictTools', () => { + it('returns the highest-probability tools, capped at the shortlist', async () => { + const { classifier, requests } = stubClient([ + { alpha: 0.6, beta: 0.2, gamma: 0.12, delta: 0.05, [NO_MATCH]: 0.03 }, + ]); + + const result = await predictTools({ + classifier, + candidates: candidates('alpha', 'beta', 'gamma', 'delta'), + request: 'find the ranked keywords', + config: CONFIG, + }); + + expect(result.names).toEqual(['alpha', 'beta', 'gamma']); + expect(result.requests).toBe(1); + expect(result.usage.inputTokens).toBe(100); + expect(requests[0].questions.needs_tool).toBeDefined(); + }); + + it('never surfaces the no-match outcome as a tool', async () => { + const { classifier } = stubClient([{ alpha: 0.3, [NO_MATCH]: 0.7 }]); + + const result = await predictTools({ + classifier, + candidates: candidates('alpha'), + request: 'do the thing', + config: CONFIG, + }); + + expect(result.names).toEqual(['alpha']); + }); + + it('drops tools below the probability floor', async () => { + const { classifier } = stubClient([{ alpha: 0.9, beta: 0.02, [NO_MATCH]: 0.08 }]); + + const result = await predictTools({ + classifier, + candidates: candidates('alpha', 'beta'), + request: 'do the thing', + config: CONFIG, + }); + + expect(result.names).toEqual(['alpha']); + }); + + it('surfaces nothing when the turn does not need a tool', async () => { + const { classifier } = stubClient([{ alpha: 0.9, [NO_MATCH]: 0.1 }], 0.04); + + const result = await predictTools({ + classifier, + candidates: candidates('alpha'), + request: 'thanks, that was helpful', + config: CONFIG, + }); + + expect(result.names).toEqual([]); + expect(result.needsTool).toBeCloseTo(0.04); + }); + + it('surfaces on every turn when the gate is set to zero', async () => { + const { classifier } = stubClient([{ alpha: 0.9, [NO_MATCH]: 0.1 }], 0); + + const result = await predictTools({ + classifier, + candidates: candidates('alpha'), + request: 'thanks', + config: { ...CONFIG, needsToolThreshold: 0 }, + }); + + expect(result.names).toEqual(['alpha']); + }); + + it('ranks a large catalog in batches and pools the winners', async () => { + const many = candidates(...Array.from({ length: 5 }, (_, i) => `tool_${i}`)); + const { classifier, requests } = stubClient([ + { tool_0: 0.7, tool_1: 0.2, [NO_MATCH]: 0.1 }, + { tool_2: 0.9, tool_3: 0.06, [NO_MATCH]: 0.04 }, + { tool_4: 0.5, [NO_MATCH]: 0.5 }, + ]); + + const result = await predictTools({ + classifier, + candidates: many, + request: 'anything', + config: { ...CONFIG, maxCatalogTools: 2 }, + }); + + expect(requests).toHaveLength(3); + expect(result.requests).toBe(3); + /** Pooled across batches and re-sorted by probability. */ + expect(result.names).toEqual(['tool_2', 'tool_0', 'tool_4']); + expect(result.usage.inputTokens).toBe(300); + }); + + it('asks the needs-a-tool question only once across batches', async () => { + const many = candidates('a', 'b', 'c', 'd'); + const { classifier, requests } = stubClient([{ a: 1 }, { c: 1 }]); + + await predictTools({ + classifier, + candidates: many, + request: 'anything', + config: { ...CONFIG, maxCatalogTools: 2 }, + }); + + expect(requests.filter((r) => r.questions.needs_tool != null)).toHaveLength(1); + }); + + it('truncates long descriptions to the configured budget', async () => { + const { classifier, requests } = stubClient([{ alpha: 1 }]); + + await predictTools({ + classifier, + candidates: [{ name: 'alpha', description: 'x'.repeat(5_000) }], + request: 'anything', + config: { ...CONFIG, descriptionChars: 50 }, + }); + + const question = requests[0].questions.best_tool; + const described = (question as { criteria: Record }).criteria.alpha; + expect(described.length).toBeLessThanOrEqual(51); + }); + + it('returns nothing, and does not call out, when there are no candidates', async () => { + const { classifier, requests } = stubClient([{ alpha: 1 }]); + + const result = await predictTools({ + classifier, + candidates: [], + request: 'anything', + config: CONFIG, + }); + + expect(result.names).toEqual([]); + expect(requests).toHaveLength(0); + }); + + it('returns nothing for an empty request', async () => { + const { classifier, requests } = stubClient([{ alpha: 1 }]); + + const result = await predictTools({ + classifier, + candidates: candidates('alpha'), + request: ' ', + config: CONFIG, + }); + + expect(result.names).toEqual([]); + expect(requests).toHaveLength(0); + }); + + it('degrades to no prediction when the judgment fails', async () => { + const classifier: Classifier = { + id: 'stub', + model: 'stub-1', + classify: async () => { + throw new Error('upstream exploded'); + }, + }; + + const result = await predictTools({ + classifier, + candidates: candidates('alpha'), + request: 'anything', + config: CONFIG, + }); + + expect(result.names).toEqual([]); + }); +}); + +describe('deferredCandidates', () => { + function registry(...tools: LCTool[]): LCToolRegistry { + return new Map(tools.map((tool) => [tool.name, tool])); + } + + it('returns only deferred tools, with their descriptions', () => { + const result = deferredCandidates( + registry( + { name: 'loaded', description: 'already here', defer_loading: false }, + { name: 'hidden', description: 'behind a search', defer_loading: true }, + ), + ); + + expect(result).toEqual([{ name: 'hidden', description: 'behind a search' }]); + }); + + it('skips a tool the model already discovered this conversation', () => { + const result = deferredCandidates( + registry({ name: 'hidden', defer_loading: true }, { name: 'found', defer_loading: true }), + new Set(['found']), + ); + + expect(result.map((c) => c.name)).toEqual(['hidden']); + }); + + it('handles a missing registry', () => { + expect(deferredCandidates(undefined)).toEqual([]); + }); +}); + +describe('namedInRequest', () => { + const catalog = candidates( + 'dataforseo_labs_google_ranked_keywords_mcp_AP11seo', + 'push_flex_message_mcp_linebot', + 'geocode_mcp_Maps', + ); + + it('finds a tool the user names by its base name', () => { + expect(namedInRequest(catalog, 'no, use push_flex_message instead')).toEqual([ + 'push_flex_message_mcp_linebot', + ]); + }); + + it('ignores case', () => { + expect(namedInRequest(catalog, 'Use GEOCODE please')).toEqual(['geocode_mcp_Maps']); + }); + + it('will not match inside a longer word', () => { + expect(namedInRequest(catalog, 'the geocoded address was wrong')).toEqual([]); + }); + + it('returns nothing for a request that names no tool', () => { + expect(namedInRequest(catalog, 'what is the weather today')).toEqual([]); + }); + + it('returns nothing for an empty request', () => { + expect(namedInRequest(catalog, '')).toEqual([]); + }); + + it('keeps looking past a match inside a longer word', () => { + expect(namedInRequest(catalog, 'geocoded badly, geocode it again')).toEqual([ + 'geocode_mcp_Maps', + ]); + }); + + it('stops at the limit when a common word matches tools on many servers', () => { + const shared = candidates('search_mcp_A', 'search_mcp_B', 'search_mcp_C', 'search_mcp_D'); + + expect(namedInRequest(shared, 'search for it', 2)).toEqual(['search_mcp_A', 'search_mcp_B']); + }); +}); + +describe('shortlistSize', () => { + const widening = { ...CONFIG, shortlist: 3, lowConfidenceExtra: 3, lowConfidenceBelow: 0.5 }; + + it('keeps the configured size when the ranking is confident', () => { + expect(shortlistSize(0.9, widening)).toBe(3); + }); + + it('widens when the ranking is not confident', () => { + expect(shortlistSize(0.3, widening)).toBe(6); + }); + + it('treats the boundary as low confidence', () => { + expect(shortlistSize(0.5, widening)).toBe(6); + }); + + it('never widens when the extra is zero', () => { + expect(shortlistSize(0, CONFIG)).toBe(CONFIG.shortlist); + }); +}); + +describe('predictTools safety behavior', () => { + const named = { ...CONFIG, surfaceNamedTools: true }; + + it('surfaces a tool the request names even when the ranking misses it', async () => { + const { classifier } = stubClient([{ alpha: 0.95, [NO_MATCH]: 0.05 }]); + + const result = await predictTools({ + classifier, + candidates: candidates('alpha', 'zebra_tool'), + request: 'no, use zebra_tool for this', + config: named, + }); + + expect(result.names).toContain('zebra_tool'); + expect(result.named).toEqual(['zebra_tool']); + }); + + it('surfaces a named tool even when the turn looks conversational', async () => { + const { classifier } = stubClient([{ alpha: 0.9, [NO_MATCH]: 0.1 }], 0.01); + + const result = await predictTools({ + classifier, + candidates: candidates('alpha', 'zebra_tool'), + request: 'thanks, though next time use zebra_tool', + config: named, + }); + + expect(result.names).toEqual(['zebra_tool']); + }); + + it('surfaces a named tool even when the judgment fails outright', async () => { + const classifier: Classifier = { + id: 'stub', + model: 'stub-1', + classify: async () => { + throw new Error('upstream exploded'); + }, + }; + + const result = await predictTools({ + classifier, + candidates: candidates('alpha', 'zebra_tool'), + request: 'use zebra_tool', + config: named, + }); + + expect(result.names).toEqual(['zebra_tool']); + }); + + it('widens the shortlist when the ranking is unconfident', async () => { + const { classifier } = stubClient([ + { a: 0.2, b: 0.19, c: 0.18, d: 0.17, e: 0.16, [NO_MATCH]: 0.1 }, + ]); + + const result = await predictTools({ + classifier, + candidates: candidates('a', 'b', 'c', 'd', 'e'), + request: 'something ambiguous', + config: { ...CONFIG, shortlist: 2, lowConfidenceExtra: 2, lowConfidenceBelow: 0.5 }, + }); + + expect(result.names).toHaveLength(4); + }); +}); + +describe('predictTools prompt overrides', () => { + it('uses the built-in wording when nothing is configured', async () => { + const { classifier, requests } = stubClient([{ alpha: 1 }]); + + await predictTools({ + classifier, + candidates: candidates('alpha'), + request: 'anything', + config: CONFIG, + }); + + const ranking = requests[0].questions.best_tool as unknown as { + instructions: { question: string; guidance: string }; + }; + expect(ranking.instructions.question).toBe(RANKING_QUESTION.instructions); + expect(ranking.instructions.guidance).toBe(RANKING_QUESTION.guidance); + expect( + (requests[0].questions.needs_tool as unknown as { instructions: string }).instructions, + ).toBe(NEEDS_TOOL_QUESTION.instructions); + }); + + it('sends operator wording instead when it is configured', async () => { + const { classifier, requests } = stubClient([{ alpha: 1 }]); + + await predictTools({ + classifier, + candidates: candidates('alpha'), + request: 'anything', + config: { + ...CONFIG, + instructions: 'Pick the tool for this support ticket.', + guidance: 'Prefer billing tools for anything about an invoice.', + needsToolInstructions: 'Does this ticket need a tool?', + }, + }); + + const ranking = requests[0].questions.best_tool as unknown as { + instructions: { question: string; guidance: string }; + }; + expect(ranking.instructions.question).toBe('Pick the tool for this support ticket.'); + expect(ranking.instructions.guidance).toBe( + 'Prefer billing tools for anything about an invoice.', + ); + expect( + (requests[0].questions.needs_tool as unknown as { instructions: string }).instructions, + ).toBe('Does this ticket need a tool?'); + }); +}); + +describe('shortlistSize with unmeasured confidence', () => { + const widening = { ...CONFIG, shortlist: 3, lowConfidenceExtra: 3, lowConfidenceBelow: 0.5 }; + + it('widens when the provider cannot measure confidence', () => { + expect(shortlistSize(null, widening)).toBe(6); + }); + + it('still respects a zero extra', () => { + expect(shortlistSize(null, CONFIG)).toBe(CONFIG.shortlist); + }); +}); + +describe('predictToolsForTurn usage', () => { + afterEach(() => { + jest.restoreAllMocks(); + }); + + it('reports what the ranking cost, priced by the configured model', async () => { + jest.spyOn(globalThis, 'fetch').mockResolvedValue( + new Response( + JSON.stringify({ + model: 'jev-1.13.0', + answers: { + best_tool: { + type: 'choice', + choice: 'search_mcp_docs', + confidence: 0.9, + probabilities: { search_mcp_docs: 0.9 }, + }, + needs_tool: { type: 'noul', noul: 0.9 }, + }, + usage: { input_tokens: 512, output_tokens: 0 }, + }), + ), + ); + const onUsage = jest.fn(); + const registry: LCToolRegistry = new Map([ + [ + 'search_mcp_docs', + { name: 'search_mcp_docs', description: 'Search the docs', defer_loading: true } as LCTool, + ], + ]); + + const names = await predictToolsForTurn({ + config: classificationSchema.parse({ + enabled: true, + provider: 'typesafe', + toolSelection: { enabled: true }, + }), + agents: [{ id: 'agent', toolRegistry: registry }], + messages: [new HumanMessage('search the docs for pricing')], + apiKey: 'test-key', + onUsage, + }); + + expect(names).toEqual(['search_mcp_docs']); + expect(onUsage).toHaveBeenCalledWith({ inputTokens: 512, outputTokens: 0 }, 'jev-latest'); + }); +}); diff --git a/packages/api/src/tools/predict.ts b/packages/api/src/tools/predict.ts new file mode 100644 index 00000000000..84167c7d0a7 --- /dev/null +++ b/packages/api/src/tools/predict.ts @@ -0,0 +1,425 @@ +import { logger } from '@librechat/data-schemas'; +import type { BaseMessage } from '@librechat/agents/langchain/messages'; +import type { TClassificationConfig } from 'librechat-data-provider'; +import type { LCToolRegistry } from '@librechat/agents'; +import type { + Classifier, + ChoiceQuestion, + BooleanQuestion, + ClassificationUsage, +} from '~/classification/types'; +import { isBooleanAnswer, isChoiceAnswer } from '~/classification/types'; +import { boolean, choice, ranked } from '~/classification/questions'; +import { classificationCapability } from '~/classification/resolve'; +import { extractDiscoveredToolsFromHistory } from '~/agents/run'; + +/** Option key standing for "none of these fit". Not a legal MCP tool name. */ +const NO_MATCH = '__no_tool_fits__'; + +/** Largest option set one ranking question may carry, including no-match. */ +const MAX_QUESTION_OPTIONS = 255; + +export const RANKING_QUESTION: { + instructions: string; + guidance: string; + noMatch: string; +} = { + instructions: + 'Which tool should the assistant call first to carry out `request`? Each option is a tool ' + + 'name; the text beside it is what that tool does.', + guidance: + 'Choose the tool whose own purpose matches the request, not one that merely shares words ' + + 'with it. Choose the no-match option when the request is conversational, is answerable ' + + 'from the conversation, or when no listed tool does this kind of work.', + noMatch: 'None of these tools would help with this request.', +}; + +export const NEEDS_TOOL_QUESTION: BooleanQuestion = boolean( + 'Carrying out `request` requires calling a tool, rather than answering from general ' + + 'knowledge or from what the conversation already contains.', + { + true: 'The request asks for an action, or for current or private data the assistant cannot already have.', + false: + 'The request is conversational, or answerable from general knowledge or the conversation so far.', + }, +); + +export interface PredictCandidate { + name: string; + description?: string; +} + +export type ToolSelectionConfig = TClassificationConfig['toolSelection']; + +export interface PredictToolsParams { + classifier: Classifier; + candidates: readonly PredictCandidate[]; + request: string; + config: ToolSelectionConfig; + signal?: AbortSignal; +} + +export interface PredictToolsResult { + names: string[]; + needsTool: number; + /** Surfaced because the request named them, not because of rank. */ + named: string[]; + usage: ClassificationUsage; + requests: number; +} + +const EMPTY_RESULT: PredictToolsResult = { + names: [], + needsTool: 0, + named: [], + usage: { inputTokens: 0, outputTokens: 0 }, + requests: 0, +}; + +function summarize(description: string | undefined, limit: number): string | null { + if (!description) { + return null; + } + const flat = description.replace(/\s+/g, ' ').trim(); + if (flat.length === 0) { + return null; + } + return flat.length > limit ? `${flat.slice(0, limit)}…` : flat; +} + +export function baseToolName(name: string): string { + const index = name.indexOf('_mcp_'); + return index === -1 ? name : name.slice(0, index); +} + +function mentionsWord(haystack: string, word: string): boolean { + for (let at = haystack.indexOf(word); at !== -1; at = haystack.indexOf(word, at + 1)) { + const before = at === 0 ? '' : haystack[at - 1]; + const after = haystack[at + word.length] ?? ''; + if (!/[a-z0-9]/.test(before) && !/[a-z0-9]/.test(after)) { + return true; + } + } + return false; +} + +/** + * Word-boundary matched, so a short name cannot match inside another word, and + * capped at `limit` so a common word ("search") cannot surface a tool from every server. + */ +export function namedInRequest( + candidates: readonly PredictCandidate[], + request: string, + limit: number = Number.POSITIVE_INFINITY, +): string[] { + const haystack = request.toLowerCase(); + if (haystack.length === 0 || limit <= 0) { + return []; + } + const found: string[] = []; + for (const candidate of candidates) { + const base = baseToolName(candidate.name).toLowerCase(); + if (base.length < 4 || !mentionsWord(haystack, base)) { + continue; + } + found.push(candidate.name); + if (found.length >= limit) { + break; + } + } + return found; +} + +export function batchCandidates( + candidates: readonly PredictCandidate[], + maxPerBatch: number, +): PredictCandidate[][] { + /** One option is spent on the no-match outcome. */ + const size = Math.max(1, Math.min(maxPerBatch, MAX_QUESTION_OPTIONS - 1)); + if (candidates.length <= size) { + return [[...candidates]]; + } + const batches: PredictCandidate[][] = []; + for (let i = 0; i < candidates.length; i += size) { + batches.push(candidates.slice(i, i + size)); + } + return batches; +} + +function rankingQuestion( + batch: readonly PredictCandidate[], + config: ToolSelectionConfig, +): ChoiceQuestion { + const criteria: Record = {}; + for (const candidate of batch) { + criteria[candidate.name] = summarize(candidate.description, config.descriptionChars); + } + criteria[NO_MATCH] = RANKING_QUESTION.noMatch; + return choice( + { + question: config.instructions ?? RANKING_QUESTION.instructions, + guidance: config.guidance ?? RANKING_QUESTION.guidance, + }, + criteria, + ); +} + +function addUsage(total: ClassificationUsage, next: ClassificationUsage): ClassificationUsage { + return { + inputTokens: total.inputTokens + (next.inputTokens ?? 0), + outputTokens: total.outputTokens + (next.outputTokens ?? 0), + }; +} + +/** + * An omission costs a round trip, an extra only tokens, so unsure widens. + * Unmeasured confidence counts as unsure for the same reason. + */ +export function shortlistSize(confidence: number | null, config: ToolSelectionConfig): number { + if (config.lowConfidenceExtra <= 0) { + return config.shortlist; + } + if (confidence != null && confidence > config.lowConfidenceBelow) { + return config.shortlist; + } + return config.shortlist + config.lowConfidenceExtra; +} + +export async function predictTools(params: PredictToolsParams): Promise { + const { classifier, candidates, request, config, signal } = params; + + const text = request?.trim(); + if (!text || candidates.length === 0) { + return EMPTY_RESULT; + } + + const named = config.surfaceNamedTools ? namedInRequest(candidates, text, config.shortlist) : []; + + const batches = batchCandidates(candidates, config.maxCatalogTools); + let usage: ClassificationUsage = { inputTokens: 0, outputTokens: 0 }; + let requests = 0; + + try { + const responses = await Promise.all( + batches.map((batch, index) => + classifier.classify({ + label: `tool-selection[${index + 1}/${batches.length}]`, + signal, + timeoutMs: config.timeoutMs, + state: { request: text }, + questions: + index === 0 + ? { + best_tool: rankingQuestion(batch, config), + needs_tool: boolean( + config.needsToolInstructions ?? NEEDS_TOOL_QUESTION.instructions, + NEEDS_TOOL_QUESTION.criteria, + ), + } + : { best_tool: rankingQuestion(batch, config) }, + }), + ), + ); + + requests = responses.length; + let needsTool = 1; + let lowestConfidence: number | null = 1; + const pooled: Array<{ name: string; probability: number }> = []; + + for (const response of responses) { + usage = addUsage(usage, response.usage); + const needs = response.answers.needs_tool; + if (isBooleanAnswer(needs)) { + needsTool = needs.probability; + } + const best = response.answers.best_tool; + if (!isChoiceAnswer(best)) { + continue; + } + if (best.confidence == null || lowestConfidence == null) { + lowestConfidence = null; + } else { + lowestConfidence = Math.min(lowestConfidence, best.confidence); + } + for (const name of ranked(best, config.minProbability)) { + if (name !== NO_MATCH) { + pooled.push({ name, probability: best.probabilities[name] }); + } + } + } + + if (needsTool < config.needsToolThreshold) { + logger.debug( + `[predictTools] needs_tool ${needsTool.toFixed(2)} below ${config.needsToolThreshold}: ` + + `surfacing ${named.length} named tool(s) only`, + ); + /** A tool the user named outright still applies; the ranker's guess does not. */ + return { names: [...named], needsTool, named, usage, requests }; + } + + const limit = shortlistSize(lowestConfidence, config); + const byProbability = pooled + .sort((a, b) => b.probability - a.probability) + .map((entry) => entry.name); + + /** Named tools do not consume the ranked shortlist's budget. */ + const names: string[] = []; + for (const name of named) { + if (!names.includes(name)) { + names.push(name); + } + } + for (const name of byProbability) { + if (names.length >= limit + named.length) { + break; + } + if (!names.includes(name)) { + names.push(name); + } + } + + logger.debug( + `[predictTools] surfaced ${names.length} of ${candidates.length} tools ` + + `in ${requests} request(s), ${usage.inputTokens} input tokens` + + (named.length > 0 ? `, ${named.length} named by the request` : '') + + (limit > config.shortlist + ? ` (widened: confidence ${lowestConfidence?.toFixed(2) ?? 'unmeasured'})` + : ''), + ); + + return { names, needsTool, named, usage, requests }; + } catch (error) { + logger.warn( + '[predictTools] prediction failed, continuing without it: ' + + (error instanceof Error ? error.message : String(error)), + ); + return { ...EMPTY_RESULT, names: named, named, usage, requests }; + } +} + +export function deferredCandidates( + registry: LCToolRegistry | undefined, + alreadyLoaded?: ReadonlySet, +): PredictCandidate[] { + if (registry == null) { + return []; + } + const candidates: PredictCandidate[] = []; + for (const tool of registry.values()) { + if (tool.defer_loading !== true) { + continue; + } + if (alreadyLoaded?.has(tool.name) === true) { + continue; + } + candidates.push({ name: tool.name, description: tool.description }); + } + return candidates; +} + +export function latestRequestText(messages: readonly BaseMessage[] | undefined): string { + if (messages == null) { + return ''; + } + for (let i = messages.length - 1; i >= 0; i--) { + const message = messages[i]; + if (message._getType?.() !== 'human') { + continue; + } + const content = message.content; + if (typeof content === 'string') { + return content.trim(); + } + if (!Array.isArray(content)) { + continue; + } + const parts: string[] = []; + for (const part of content) { + if (typeof part === 'string') { + parts.push(part); + } else if (part != null && typeof part === 'object' && 'text' in part) { + const text = (part as { text?: unknown }).text; + if (typeof text === 'string') { + parts.push(text); + } + } + } + return parts.join(' ').trim(); + } + return ''; +} + +export interface AgentWithRegistry { + id?: string; + toolRegistry?: LCToolRegistry; +} + +export interface PredictToolsForTurnParams { + config?: TClassificationConfig | null; + agents: readonly AgentWithRegistry[]; + messages: readonly BaseMessage[] | undefined; + signal?: AbortSignal; + apiKey?: string; + alreadyLoaded?: ReadonlySet; + /** Receives what the ranking cost, so the caller can bill it. */ + onUsage?: (usage: ClassificationUsage, model: string) => void; +} + +export async function predictToolsForTurn(params: PredictToolsForTurnParams): Promise { + const capability = classificationCapability(params.config, 'toolSelection', { + apiKey: params.apiKey, + }); + if (capability == null) { + logger.debug('[predictToolsForTurn] skipped: tool selection is off, or no usable classifier'); + return []; + } + + const alreadyLoaded = + params.alreadyLoaded ?? + (params.messages?.length ? extractDiscoveredToolsFromHistory(params.messages) : undefined); + + /** Names are unique across agents; `createRun` routes each to its owner. */ + const seen = new Set(); + const candidates: PredictCandidate[] = []; + for (const agent of params.agents) { + for (const candidate of deferredCandidates(agent.toolRegistry, alreadyLoaded)) { + if (seen.has(candidate.name)) { + continue; + } + seen.add(candidate.name); + candidates.push(candidate); + } + } + + if (candidates.length === 0) { + logger.debug( + `[predictToolsForTurn] skipped: none of the ${params.agents.length} agent(s) defer a tool`, + ); + return []; + } + + const request = latestRequestText(params.messages); + logger.debug( + `[predictToolsForTurn] ranking ${candidates.length} deferred tool(s) against ` + + `${request.length} characters of request`, + ); + + const result = await predictTools({ + classifier: capability.classifier, + config: capability.settings, + candidates, + request, + signal: params.signal, + }); + + if (result.requests > 0) { + params.onUsage?.(result.usage, capability.classifier.model); + } + + logger.debug( + `[predictToolsForTurn] surfacing ${result.names.length} tool(s)` + + (result.names.length > 0 ? `: ${result.names.join(', ')}` : ''), + ); + + return result.names; +} diff --git a/packages/data-provider/src/config.ts b/packages/data-provider/src/config.ts index 969b3305134..89a4c6244b0 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2908,12 +2908,95 @@ export type TOpenIdDiscoveryConfig = z.infer; /** Maximum CAS attempts per ACL document, including the initial attempt. */ export const permissionWriteAttemptsSchema = z.number().int().min(1).max(100).default(3); +export const classificationProviderSchema = z + .object({ + baseURL: z.string().url().optional(), + model: z.string().optional(), + /** Which wire vocabulary the endpoint speaks. */ + dialect: z.enum(['port', 'systemone']).optional(), + /** Nests `state` and `questions` under this key, for hosts that wrap them. */ + requestKey: z.string().optional(), + /** Reads the answer envelope from this key, for hosts that wrap the response. */ + responseKey: z.string().optional(), + /** Ceiling for one judgment, retries and backoff included. A judgment that + * misses it is abandoned, never awaited. */ + timeoutMs: z.number().int().positive().max(60_000).optional(), + /** Retries inside `timeoutMs`, for a rate limit, server or network error only. */ + maxRetries: z.number().int().nonnegative().max(5).optional(), + /** Environment variable holding this provider's key. Never the key itself. */ + apiKeyEnv: z.string().optional(), + }) + /** Strict so a misspelled key fails loudly here rather than as a missing + * setting much later. */ + .strict(); + +export type TClassificationProviderConfig = z.infer; + +/** Every capability defaults to off, so an unset block changes nothing. */ +export const classificationSchema = z.object({ + enabled: z.boolean().default(false), + /** Which registered provider answers. An unknown name disables classification. */ + provider: z.string().default('http'), + providers: z.record(z.string(), classificationProviderSchema).default({}), + /** + * Surfaces deferred tools the turn is likely to need. Only ever adds: a tool + * it passes over stays listed by name and one `tool_search` away. + */ + toolSelection: z + .object({ + enabled: z.boolean().default(false), + /** Overrides the provider's timeout. Ranking a large catalog is a much + * bigger request than a yes/no question and can need longer. */ + timeoutMs: z.number().int().positive().max(60_000).optional(), + shortlist: z.number().int().positive().max(50).default(5), + minProbability: z.number().min(0).max(1).default(0.05), + /** Below this, only tools the request names outright are surfaced. */ + needsToolThreshold: z.number().min(0).max(1).default(0.15), + /** An omission costs a round trip, so an unsure ranking widens instead. */ + lowConfidenceExtra: z.number().int().nonnegative().max(20).default(3), + lowConfidenceBelow: z.number().min(0).max(1).default(0.5), + /** Covers "no, use the other tool" without waiting for the next ranking. */ + surfaceNamedTools: z.boolean().default(true), + /** Above this the catalog is ranked in batches. */ + maxCatalogTools: z.number().int().positive().max(250).default(200), + descriptionChars: z.number().int().positive().max(2_000).default(300), + /** Replaces the ranking question. Unset uses the built-in wording. */ + instructions: z.string().min(1).max(4_000).optional(), + /** Replaces the rubric shown beside the ranking question. */ + guidance: z.string().min(1).max(4_000).optional(), + /** Replaces the question asking whether the turn needs a tool at all. */ + needsToolInstructions: z.string().min(1).max(4_000).optional(), + }) + .default({}), + /** Skips the memory model on turns that do not ask to remember, update or forget anything. */ + memoryGate: z + .object({ + enabled: z.boolean().default(false), + timeoutMs: z.number().int().positive().max(60_000).optional(), + threshold: z.number().min(0).max(1).default(0.5), + instructions: z.string().min(1).max(4_000).optional(), + /** + * What a yes and a no mean. Named `whenTrue`/`whenFalse` because YAML + * reads bare `true:` and `false:` keys as booleans, not strings. + */ + whenTrue: z.string().min(1).max(4_000).optional(), + whenFalse: z.string().min(1).max(4_000).optional(), + categorize: z.boolean().default(false), + categoryThreshold: z.number().min(0).max(1).default(0.4), + detectUpdates: z.boolean().default(false), + }) + .default({}), +}); + +export type TClassificationConfig = z.infer; + export const configSchema = z.object({ version: z.string(), permissions: z.object({ maxWriteAttempts: permissionWriteAttemptsSchema }).optional(), cache: z.boolean().default(true), ocr: ocrSchema.optional(), webSearch: webSearchSchema.optional(), + classification: classificationSchema.optional(), langfuse: langfuseConfigSchema.optional(), memory: memorySchema.optional(), summarization: summarizationConfigSchema.optional(), diff --git a/packages/data-schemas/src/app/service.spec.ts b/packages/data-schemas/src/app/service.spec.ts index 820335f4b5d..a5e330f92ee 100644 --- a/packages/data-schemas/src/app/service.spec.ts +++ b/packages/data-schemas/src/app/service.spec.ts @@ -4,7 +4,12 @@ import { defaultAssistantsVersion, } from 'librechat-data-provider'; import type { DeepPartial, TCustomConfig } from 'librechat-data-provider'; -import { AppService, loadFiltersConfig, loadSummarizationConfig } from './service'; +import { + AppService, + loadFiltersConfig, + loadSummarizationConfig, + loadClassificationConfig, +} from './service'; import logger from '~/config/winston'; jest.mock('~/config/winston', () => ({ @@ -84,6 +89,31 @@ describe('loadSummarizationConfig', () => { }); }); +describe('loadClassificationConfig', () => { + it('is null when the block is absent', () => { + expect(loadClassificationConfig({})).toBeNull(); + }); + + it('fills every schema default the yaml leaves out', () => { + const config = loadClassificationConfig({ + classification: { + enabled: true, + toolSelection: { enabled: true }, + memoryGate: { enabled: true }, + }, + }); + + expect(config?.toolSelection.maxCatalogTools).toBe(200); + expect(config?.toolSelection.minProbability).toBe(0.05); + expect(config?.memoryGate.categoryThreshold).toBe(0.4); + expect(config?.memoryGate.threshold).toBe(0.5); + }); + + it('turns classification off rather than running on an invalid block', () => { + expect(loadClassificationConfig({ classification: { enabled: 'yes' } } as never)).toBeNull(); + }); +}); + describe('loadFiltersConfig', () => { it('treats omission and zero-rule source configs as disabled', () => { expect(loadFiltersConfig({})).toBeUndefined(); diff --git a/packages/data-schemas/src/app/service.ts b/packages/data-schemas/src/app/service.ts index e762d1620e8..2c5f9adabd2 100644 --- a/packages/data-schemas/src/app/service.ts +++ b/packages/data-schemas/src/app/service.ts @@ -1,6 +1,7 @@ import { AgentCapabilities, EModelEndpoint, + classificationSchema, filtersConfigSchema, hasActiveFiltersConfig, getConfigDefaults, @@ -76,6 +77,24 @@ export function loadSkillSyncConfig(config: DeepPartial): AppConf return parsed.data; } +/** The loaded yaml is unparsed, so schema defaults only exist after this parse. */ +export function loadClassificationConfig( + config: DeepPartial, +): AppConfig['classification'] { + const raw = config.classification; + if (!raw || typeof raw !== 'object') { + return null; + } + + const parsed = classificationSchema.safeParse(raw); + if (!parsed.success) { + logger.warn('[AppService] Invalid classification config', parsed.error.flatten()); + return null; + } + + return parsed.data; +} + export function loadLangfuseConfig(config: DeepPartial): AppConfig['langfuse'] { const raw = config.langfuse; if (!raw || typeof raw !== 'object') { @@ -161,6 +180,7 @@ export const AppService = async (params?: { const mcpServersConfig = config.mcpServers || null; const mcpSettings = config.mcpSettings || null; + const classification = loadClassificationConfig(config); const actions = config.actions; const registration = config.registration ?? configDefaults.registration; const interfaceConfig = await loadDefaultInterface({ config, configDefaults }); @@ -181,6 +201,7 @@ export const AppService = async (params?: { skillSync, webSearch, mcpSettings, + classification, fileStrategy, registration, transactions, diff --git a/packages/data-schemas/src/methods/tx.spec.ts b/packages/data-schemas/src/methods/tx.spec.ts index 1e22c0cd0f4..b9b422cb8ec 100644 --- a/packages/data-schemas/src/methods/tx.spec.ts +++ b/packages/data-schemas/src/methods/tx.spec.ts @@ -254,6 +254,14 @@ describe('getValueKey', () => { expect(getValueKey('gpt-oss-20b')).toBe('gpt-oss-20b'); expect(getValueKey('oai/gpt-oss:20b')).toBe('gpt-oss:20b'); }); + + it('prices every name the Jev classifier is served under', () => { + for (const model of ['jev-latest', '~typesafe/jev-latest', 'typesafe/jev', 'jev-1.13.0']) { + expect(getValueKey(model)).toBe('jev'); + } + expect(getMultiplier({ valueKey: 'jev', tokenType: 'prompt' })).toBe(0.042); + expect(getMultiplier({ valueKey: 'jev', tokenType: 'completion' })).toBe(0); + }); }); describe('getMultiplier', () => { diff --git a/packages/data-schemas/src/methods/tx.ts b/packages/data-schemas/src/methods/tx.ts index 7661677b4ce..f4ca7e4dea6 100644 --- a/packages/data-schemas/src/methods/tx.ts +++ b/packages/data-schemas/src/methods/tx.ts @@ -125,6 +125,8 @@ export const tokenValues: Record deepseek: { prompt: 0.28, completion: 0.42 }, command: { prompt: 0.38, completion: 0.38 }, gemma: { prompt: 0.02, completion: 0.04 }, + /** TypeSafe's classifier; output tokens are free. */ + jev: { prompt: 0.042, completion: 0 }, gemini: { prompt: 0.5, completion: 1.5 }, 'gpt-oss': { prompt: 0.05, completion: 0.2 }, 'gpt-3.5-turbo-1106': { prompt: 1, completion: 2 }, diff --git a/packages/data-schemas/src/types/app.ts b/packages/data-schemas/src/types/app.ts index df2a4b60540..b733d5e0465 100644 --- a/packages/data-schemas/src/types/app.ts +++ b/packages/data-schemas/src/types/app.ts @@ -15,6 +15,7 @@ import type { SummarizationConfig, SkillSyncConfig, FiltersConfig, + TClassificationConfig, } from 'librechat-data-provider'; export type JsonSchemaType = { @@ -102,6 +103,8 @@ export interface AppConfig { mcpConfig?: TCustomConfig['mcpServers'] | null; /** MCP settings (domain allowlist, etc.) */ mcpSettings?: TCustomConfig['mcpSettings'] | null; + /** Classification provider and the capabilities that consult it */ + classification?: TClassificationConfig | null; /** File configuration */ fileConfig?: TFileConfig; /** Secure image links configuration, enabled unless explicitly disabled */