From 84fbaee28826c21cf862ffeb8d2007d5fce70555 Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Tue, 22 Sep 2026 21:25:46 +0900 Subject: [PATCH 01/11] feat: add a classification provider interface --- .env.example | 14 + librechat.example.yaml | 33 ++ .../api/src/classification/config.spec.ts | 134 ++++++++ packages/api/src/classification/index.ts | 5 + .../api/src/classification/presets.spec.ts | 120 +++++++ .../classification/providers/dialect.spec.ts | 65 ++++ .../src/classification/providers/dialect.ts | 49 +++ .../src/classification/providers/http.spec.ts | 295 ++++++++++++++++++ .../api/src/classification/providers/http.ts | 112 +++++++ .../src/classification/providers/transport.ts | 191 ++++++++++++ .../api/src/classification/questions.spec.ts | 124 ++++++++ packages/api/src/classification/questions.ts | 82 +++++ packages/api/src/classification/registry.ts | 87 ++++++ packages/api/src/classification/resolve.ts | 89 ++++++ packages/api/src/classification/types.ts | 128 ++++++++ packages/api/src/index.ts | 2 + packages/data-provider/src/config.ts | 34 ++ packages/data-schemas/src/app/service.ts | 2 + packages/data-schemas/src/types/app.ts | 2 + 19 files changed, 1568 insertions(+) create mode 100644 packages/api/src/classification/config.spec.ts create mode 100644 packages/api/src/classification/index.ts create mode 100644 packages/api/src/classification/presets.spec.ts create mode 100644 packages/api/src/classification/providers/dialect.spec.ts create mode 100644 packages/api/src/classification/providers/dialect.ts create mode 100644 packages/api/src/classification/providers/http.spec.ts create mode 100644 packages/api/src/classification/providers/http.ts create mode 100644 packages/api/src/classification/providers/transport.ts create mode 100644 packages/api/src/classification/questions.spec.ts create mode 100644 packages/api/src/classification/questions.ts create mode 100644 packages/api/src/classification/registry.ts create mode 100644 packages/api/src/classification/resolve.ts create mode 100644 packages/api/src/classification/types.ts 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/librechat.example.yaml b/librechat.example.yaml index 260912aa780..5d159a89ef7 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -515,6 +515,39 @@ actions: # - 'host.docker.internal:8080' # - '127.0.0.1:8080' +# Classification: small typed judgments (yes/no, pick one, rate) that code can +# branch on. Off unless enabled. The key is read from the environment variable +# named by apiKeyEnv, never written here. +# 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 + # Example MCP Servers Object Structure # mcpServers: # everything: diff --git a/packages/api/src/classification/config.spec.ts b/packages/api/src/classification/config.spec.ts new file mode 100644 index 00000000000..b1421305e79 --- /dev/null +++ b/packages/api/src/classification/config.spec.ts @@ -0,0 +1,134 @@ +import { classificationSchema } from 'librechat-data-provider'; +import type { TClassificationConfig } from 'librechat-data-provider'; +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 = async (url: string, init: { body?: string }) => { + calls.push({ url, body: JSON.parse(init.body ?? '{}') }); + return { + ok: true, + status: 200, + headers: { get: () => null }, + text: async () => JSON.stringify(response), + }; + }; + return { calls, fetch: fetch as never }; +} + +function parse(raw: unknown): TClassificationConfig { + return classificationSchema.parse(raw) as TClassificationConfig; +} + +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..0bfcd642ad4 --- /dev/null +++ b/packages/api/src/classification/presets.spec.ts @@ -0,0 +1,120 @@ +import type { TClassificationConfig } from 'librechat-data-provider'; +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 = async (url: string, init: { body?: string }) => { + calls.push({ url, body: JSON.parse(init.body ?? '{}') }); + return { + ok: true, + status: 200, + headers: { get: () => null }, + text: async () => JSON.stringify(response), + }; + }; + return { calls, fetch: fetch as never }; +} + +function configFor(provider: string, settings?: Record): TClassificationConfig { + return { + enabled: true, + provider, + providers: settings == null ? {} : { [provider]: settings }, + toolSelection: {}, + memoryGate: {}, + } as unknown as TClassificationConfig; +} + +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..eaeb4c937ab --- /dev/null +++ b/packages/api/src/classification/providers/dialect.spec.ts @@ -0,0 +1,65 @@ +import { toWireQuestion, readAnswer } from './dialect'; +import { boolean, choice, score } 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 choice and score alone in either dialect', () => { + const pick = choice('which', { a: null, b: null }); + const rate = score('how much', ['low', 'high']); + + expect(toWireQuestion(pick, 'systemone').type).toBe('choice'); + expect(toWireQuestion(rate, 'systemone').type).toBe('score'); + 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..1f267a290b0 --- /dev/null +++ b/packages/api/src/classification/providers/dialect.ts @@ -0,0 +1,49 @@ +import type { ClassificationAnswer, ClassificationQuestion } from '../types'; + +export type Dialect = 'port' | 'systemone'; + +interface WireQuestion { + type: string; + instructions: unknown; + criteria?: unknown; +} + +interface WireAnswer { + type?: unknown; + probability?: unknown; + noul?: unknown; + choice?: unknown; + score?: 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 }; + } + if (record.type === 'score' && typeof record.score === 'number') { + return { type: 'score', score: record.score, 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..6234637ed4e --- /dev/null +++ b/packages/api/src/classification/providers/http.spec.ts @@ -0,0 +1,295 @@ +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 choice and score answers through', 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 } }, + rate: { type: 'score', instructions: 'how much', criteria: ['low', 'high'] }, + }, + }); + + expect(result.answers.pick).toMatchObject({ type: 'choice', choice: 'b', confidence: 0.7 }); + expect(result.answers.rate).toMatchObject({ type: 'score', score: 1.4 }); + }); + + 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, + sleep: async (ms) => { + waits.push(ms); + }, + }); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(waits).toEqual([10_000]); + }); + + 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..407ad65adad --- /dev/null +++ b/packages/api/src/classification/providers/http.ts @@ -0,0 +1,112 @@ +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 { + 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 send = createTransport({ providerId: PROVIDER_ID, ...options }); + const model = options.model ?? ''; + const dialect: Dialect = options.dialect ?? 'port'; + const { requestKey, responseKey } = options; + + return { + id: PROVIDER_ID, + 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, PROVIDER_ID, (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..c18a5302cca --- /dev/null +++ b/packages/api/src/classification/providers/transport.ts @@ -0,0 +1,191 @@ +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) => 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; +} + +const defaultSleep = (ms: number): Promise => + new Promise((resolve) => setTimeout(resolve, ms)); + +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); + } + } + + return async function send(payload, signal, label, timeoutOverrideMs) { + const timeoutMs = + timeoutOverrideMs != null && timeoutOverrideMs > 0 ? timeoutOverrideMs : defaultTimeoutMs; + let lastError: ClassificationError | undefined; + for (let attemptNo = 0; attemptNo <= maxRetries; attemptNo++) { + try { + const started = Date.now(); + const body = await attempt(payload, signal, timeoutMs); + 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; + } + await sleep( + lastError.retryAfterMs ?? BACKOFF_MS[Math.min(attemptNo, BACKOFF_MS.length - 1)], + ); + } + } + 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..7551b1346cc --- /dev/null +++ b/packages/api/src/classification/questions.spec.ts @@ -0,0 +1,124 @@ +import type { ScoreAnswer, ChoiceAnswer, BooleanAnswer } from './types'; +import { + score, + level, + label, + choice, + isTrue, + ranked, + boolean, + normalized, + probabilityOf, +} 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 and a score', () => { + expect(choice('Which team?', { billing: 'money', tech: 'bugs' })).toEqual({ + type: 'choice', + instructions: 'Which team?', + criteria: { billing: 'money', tech: 'bugs' }, + }); + expect(score('How angry?', ['Calm', 'Cross', 'Furious'])).toEqual({ + type: 'score', + instructions: 'How angry?', + criteria: ['Calm', 'Cross', 'Furious'], + }); + }); +}); + +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('reads one option probability', () => { + expect(probabilityOf(answer, 'billing')).toBeCloseTo(0.2); + }); + + it('returns zero for an option that was never offered', () => { + expect(probabilityOf(answer, 'absent')).toBe(0); + }); + + 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([]); + }); +}); + +describe('score helpers', () => { + const levels = ['Calm', 'Frustrated', 'Very angry']; + const answer: ScoreAnswer = { + type: 'score', + score: 1.24, + confidence: 0.6, + probabilities: { '0': 0.12, '1': 0.52, '2': 0.36 }, + }; + + it('reports the most probable level, not the rounded score', () => { + expect(level(answer)).toBe(1); + }); + + it('labels that level', () => { + expect(label(answer, levels)).toBe('Frustrated'); + }); + + it('returns an empty label when the levels do not cover it', () => { + expect(label(answer, ['only one'])).toBe(''); + }); + + it('normalizes the weighted score against the top level', () => { + expect(normalized(answer, levels.length)).toBeCloseTo(0.62); + }); + + it('clamps a score outside the level range', () => { + expect(normalized({ ...answer, score: 99 }, levels.length)).toBe(1); + expect(normalized({ ...answer, score: -1 }, levels.length)).toBe(0); + }); + + it('returns zero when there are too few levels to normalize', () => { + expect(normalized(answer, 1)).toBe(0); + }); + + it('handles a missing answer', () => { + expect(level(undefined)).toBe(0); + expect(normalized(undefined, 3)).toBe(0); + }); +}); diff --git a/packages/api/src/classification/questions.ts b/packages/api/src/classification/questions.ts new file mode 100644 index 00000000000..471ffe836b8 --- /dev/null +++ b/packages/api/src/classification/questions.ts @@ -0,0 +1,82 @@ +import type { + ScoreAnswer, + ChoiceAnswer, + ScoreQuestion, + 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 score( + instructions: ClassificationText, + levels: ClassificationText[], +): ScoreQuestion { + return { type: 'score', instructions, criteria: levels }; +} + +export function isTrue(answer: BooleanAnswer | undefined, threshold: number): boolean { + return answer != null && answer.probability >= threshold; +} + +export function probabilityOf( + answer: ChoiceAnswer | ScoreAnswer | undefined, + option: string, +): number { + return answer?.probabilities?.[option] ?? 0; +} + +/** 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); +} + +/** The most probable level, which is not always the rounded weighted score. */ +export function level(answer: ScoreAnswer | undefined): number { + if (answer == null) { + return 0; + } + let best = 0; + let bestProbability = -1; + for (const [key, probability] of Object.entries(answer.probabilities)) { + if (probability > bestProbability) { + bestProbability = probability; + best = Number(key); + } + } + return Number.isFinite(best) ? best : 0; +} + +export function label(answer: ScoreAnswer | undefined, levels: readonly string[]): string { + return levels[level(answer)] ?? ''; +} + +/** The weighted score as a fraction of the highest level. */ +export function normalized(answer: ScoreAnswer | undefined, levelCount: number): number { + if (answer == null || levelCount < 2) { + return 0; + } + return Math.min(Math.max(answer.score / (levelCount - 1), 0), 1); +} diff --git a/packages/api/src/classification/registry.ts b/packages/api/src/classification/registry.ts new file mode 100644 index 00000000000..65e2ec6a5ae --- /dev/null +++ b/packages/api/src/classification/registry.ts @@ -0,0 +1,87 @@ +import type { ProviderFetch } from './providers/transport'; +import type { Dialect } from './providers/dialect'; +import type { Classifier } from './types'; +import { createHttpClassifier } from './providers/http'; + +export interface ProviderSettings { + baseURL?: string; + model?: string; + dialect?: Dialect; + requestKey?: string; + responseKey?: string; + timeoutMs?: number; + maxRetries?: number; + apiKeyEnv?: string; +} + +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, +): Classifier { + return createHttpClassifier({ + 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..a5734e4ad81 --- /dev/null +++ b/packages/api/src/classification/resolve.ts @@ -0,0 +1,89 @@ +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); + 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; + } +} diff --git a/packages/api/src/classification/types.ts b/packages/api/src/classification/types.ts new file mode 100644 index 00000000000..9a3cc9ae7da --- /dev/null +++ b/packages/api/src/classification/types.ts @@ -0,0 +1,128 @@ +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 interface ScoreQuestion { + type: 'score'; + instructions: ClassificationText; + criteria: ClassificationText[]; +} + +export type ClassificationQuestion = BooleanQuestion | ChoiceQuestion | ScoreQuestion; + +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 interface ScoreAnswer { + type: 'score'; + /** 0 to levels - 1. */ + score: number; + confidence: number | null; + probabilities: Record; +} + +export type ClassificationAnswer = BooleanAnswer | ChoiceAnswer | ScoreAnswer; + +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'; +} + +export function isScoreAnswer(answer: ClassificationAnswer | undefined): answer is ScoreAnswer { + return answer?.type === 'score'; +} 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/data-provider/src/config.ts b/packages/data-provider/src/config.ts index 969b3305134..9d7645289a7 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2908,12 +2908,46 @@ 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(), + /** Per-request ceiling. A judgment that misses it is abandoned, never awaited. */ + timeoutMs: z.number().int().positive().max(60_000).optional(), + /** Retries for a rate limit or a server 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({}), +}); + +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.ts b/packages/data-schemas/src/app/service.ts index e762d1620e8..021c8763f55 100644 --- a/packages/data-schemas/src/app/service.ts +++ b/packages/data-schemas/src/app/service.ts @@ -161,6 +161,7 @@ export const AppService = async (params?: { const mcpServersConfig = config.mcpServers || null; const mcpSettings = config.mcpSettings || null; + const classification = config.classification || null; const actions = config.actions; const registration = config.registration ?? configDefaults.registration; const interfaceConfig = await loadDefaultInterface({ config, configDefaults }); @@ -181,6 +182,7 @@ export const AppService = async (params?: { skillSync, webSearch, mcpSettings, + classification, fileStrategy, registration, transactions, diff --git a/packages/data-schemas/src/types/app.ts b/packages/data-schemas/src/types/app.ts index df2a4b60540..fc188edb50e 100644 --- a/packages/data-schemas/src/types/app.ts +++ b/packages/data-schemas/src/types/app.ts @@ -102,6 +102,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?: TCustomConfig['classification'] | null; /** File configuration */ fileConfig?: TFileConfig; /** Secure image links configuration, enabled unless explicitly disabled */ From 7d3e2103ada5ca74eda384488f5aa5a0df338298 Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Tue, 22 Sep 2026 13:04:28 +0900 Subject: [PATCH 02/11] feat: select deferred tools and gate memory with a classifier --- api/server/controllers/agents/client.js | 41 ++ librechat.example.yaml | 24 ++ packages/api/src/agents/memory.ts | 12 + packages/api/src/agents/run.ts | 11 + packages/api/src/classification/resolve.ts | 21 + packages/api/src/memory/gate.spec.ts | 189 +++++++++ packages/api/src/memory/gate.ts | 131 ++++++ packages/api/src/memory/index.ts | 1 + packages/api/src/tools/index.ts | 1 + packages/api/src/tools/predict.spec.ts | 468 +++++++++++++++++++++ packages/api/src/tools/predict.ts | 401 ++++++++++++++++++ packages/data-provider/src/config.ts | 42 ++ 12 files changed, 1342 insertions(+) create mode 100644 packages/api/src/memory/gate.spec.ts create mode 100644 packages/api/src/memory/gate.ts create mode 100644 packages/api/src/tools/predict.spec.ts create mode 100644 packages/api/src/tools/predict.ts diff --git a/api/server/controllers/agents/client.js b/api/server/controllers/agents/client.js index fcf65b8c028..e00edc016e1 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, @@ -1853,6 +1856,19 @@ class AgentClient extends BaseClient { } /** Builds the independently opt-in live reasoning-label controller. */ + /** @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, + }); + } + buildReasoningLabelWiring(streamId, abortSignal, seedFromContent = false) { if (!streamId || typeof Run?.prototype?.generateReasoningLabel !== 'function') { return undefined; @@ -3383,6 +3399,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 +4396,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 +4826,18 @@ class AgentClient extends BaseClient { if (this.agentConfigs && this.agentConfigs.size > 0) { agents.push(...this.agentConfigs.values()); } + + /** Ahead of the checkpoint setup below: the prune and `createRun` are + * deliberately overlapped, and an await between them serializes both. */ + const predictedToolNames = await predictToolsForTurn({ + config: appConfig?.classification, + agents, + messages, + signal: abortController.signal, + }); + if (predictedToolNames.length > 0) { + this.predictedToolNames = predictedToolNames; + } const modelBoundCallback = AgentClient.prototype.createModelBoundChatModelCallback.call(this); const initialModelBoundAdmission = @@ -4916,6 +4956,7 @@ class AgentClient extends BaseClient { messages, discoveredToolNames: this.eventActorContinuation === 'warm' ? this.eventActorDiscoveredToolNames : undefined, + predictedToolNames, modelCallbacks: [ modelBoundCallback, createAgentMemoryCallback(this.attachmentMemoryContext ?? {}), diff --git a/librechat.example.yaml b/librechat.example.yaml index 5d159a89ef7..e9ed36980a0 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -547,6 +547,30 @@ actions: # 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 +# # 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 carry nothing durable. +# memoryGate: +# enabled: true +# threshold: 0.25 +# # whenTrue/whenFalse rather than true/false: YAML reads those bare keys +# # as booleans. +# whenTrue: A lasting preference or a fact about the user. +# whenFalse: Small talk, or a detail that only matters in this task. # Example MCP Servers Object Structure # mcpServers: diff --git a/packages/api/src/agents/memory.ts b/packages/api/src/agents/memory.ts index 9e7627d9aa2..4d5c39a48b8 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,14 @@ export async function createMemoryProcessor({ messages: BaseMessage[], inspectionMessages?: BaseMessage[], ): Promise<(TAttachment | null)[] | undefined> { + if (gate != null && !(await gate(messages))) { + logger.debug('[MemoryAgent] Turn carries nothing durable; skipping', { + userId, + conversationId, + messageId, + }); + return undefined; + } try { return await processMemory({ res, diff --git a/packages/api/src/agents/run.ts b/packages/api/src/agents/run.ts index c6287521cc7..34d2dbeea60 100644 --- a/packages/api/src/agents/run.ts +++ b/packages/api/src/agents/run.ts @@ -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/resolve.ts b/packages/api/src/classification/resolve.ts index a5734e4ad81..3964bedf41c 100644 --- a/packages/api/src/classification/resolve.ts +++ b/packages/api/src/classification/resolve.ts @@ -87,3 +87,24 @@ export function resolveClassifier(params: ResolveClassifierParams): Classifier | 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/memory/gate.spec.ts b/packages/api/src/memory/gate.spec.ts new file mode 100644 index 00000000000..bdbc977b8bd --- /dev/null +++ b/packages/api/src/memory/gate.spec.ts @@ -0,0 +1,189 @@ +import { HumanMessage, AIMessage } from '@librechat/agents/langchain/messages'; +import type { + Classifier, + ClassificationResult, + ClassificationRequest, +} from '~/classification/types'; +import type { MemoryGateSettings } from './gate'; +import { createMemoryGate, transcribeTail, DURABLE_QUESTION } from './gate'; + +const ON: MemoryGateSettings = { enabled: true, threshold: 0.25 }; + +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: { durable: { 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: { enabled: false, threshold: 0.25 } }), + ).toBeNull(); + }); + + it('processes a turn that carries something durable', async () => { + const { classifier, requests } = stubClassifier(0.88); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.(TURN)).resolves.toBe(true); + expect(requests).toHaveLength(1); + expect(requests[0].questions.durable).toBeDefined(); + }); + + it('skips a turn that carries nothing durable', async () => { + const { classifier } = stubClassifier(0.03); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.([new HumanMessage('thanks!')])).resolves.toBe(false); + }); + + it('treats the threshold as inclusive', async () => { + const { classifier } = stubClassifier(0.25); + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.(TURN)).resolves.toBe(true); + }); + + it('respects a stricter threshold', async () => { + const { classifier } = stubClassifier(0.5); + const gate = createMemoryGate({ classifier, settings: { enabled: true, threshold: 0.8 } }); + + await expect(gate?.(TURN)).resolves.toBe(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?.(TURN)).resolves.toBe(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: { durable: { type: 'choice', choice: 'yes', confidence: 1, probabilities: {} } }, + usage: { inputTokens: 1, outputTokens: 1 }, + }), + }; + const gate = createMemoryGate({ classifier, settings: ON }); + + await expect(gate?.(TURN)).resolves.toBe(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?.([])).resolves.toBe(false); + expect(requests).toHaveLength(0); + }); +}); + +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?.(TURN); + + const question = requests[0].questions.durable as unknown as { + instructions: string; + criteria: { true: string; false: string }; + }; + expect(question.instructions).toBe(DURABLE_QUESTION.instructions); + expect(question.criteria.true).toBe(DURABLE_QUESTION.criteria?.true); + expect(question.criteria.false).toBe(DURABLE_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?.(TURN); + + const question = requests[0].questions.durable 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.'); + }); +}); diff --git a/packages/api/src/memory/gate.ts b/packages/api/src/memory/gate.ts new file mode 100644 index 00000000000..0246af4f728 --- /dev/null +++ b/packages/api/src/memory/gate.ts @@ -0,0 +1,131 @@ +import { logger } from '@librechat/data-schemas'; +import type { BaseMessage } from '@librechat/agents/langchain/messages'; +import type { TClassificationConfig } from 'librechat-data-provider'; +import type { BooleanQuestion } from '~/classification/types'; +import type { Classifier } from '~/classification/types'; +import { boolean, isTrue } from '~/classification/questions'; +import { isBooleanAnswer } from '~/classification/types'; + +export type MemoryGate = (messages: BaseMessage[]) => Promise; + +export type MemoryGateSettings = TClassificationConfig['memoryGate']; + +export interface CreateMemoryGateParams { + classifier: Classifier; + settings: MemoryGateSettings; + signal?: AbortSignal; + windowSize?: number; + maxChars?: number; +} + +const DEFAULT_WINDOW = 6; +const DEFAULT_MAX_CHARS = 8_000; + +export const DURABLE_QUESTION: BooleanQuestion = boolean( + 'Does `conversation` hold something about this user that would still matter in an unrelated ' + + 'conversation weeks from now?', + { + true: + 'A lasting preference, a fact about who they are or what they work on, or a decision ' + + 'they want remembered.', + false: 'Small talk, or a detail that only matters inside this task.', + }, +); + +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'); +} + +export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | null { + const { classifier, settings, signal } = 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 question = boolean(settings.instructions ?? DURABLE_QUESTION.instructions, { + true: settings.whenTrue ?? DURABLE_QUESTION.criteria?.true, + false: settings.whenFalse ?? DURABLE_QUESTION.criteria?.false, + }); + + return async function memoryGate(messages: BaseMessage[]): Promise { + const transcript = transcribeTail(messages ?? [], windowSize, maxChars); + if (transcript.length === 0) { + return false; + } + try { + const response = await classifier.classify({ + label: 'memory-gate', + signal, + state: { conversation: transcript }, + questions: { durable: question }, + }); + const answer = response.answers.durable; + if (!isBooleanAnswer(answer)) { + return true; + } + const keep = isTrue(answer, threshold); + logger.debug( + `[memoryGate] durable ${answer.probability.toFixed(2)} vs threshold ${threshold}: ` + + `${keep ? 'processing' : 'skipping'} memory`, + ); + return keep; + } catch (error) { + logger.warn( + '[memoryGate] judgment failed, processing memory as usual: ' + + (error instanceof Error ? error.message : String(error)), + ); + return true; + } + }; +} 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..f8b8318e2c4 --- /dev/null +++ b/packages/api/src/tools/predict.spec.ts @@ -0,0 +1,468 @@ +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, + 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([]); + }); +}); + +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); + }); +}); diff --git a/packages/api/src/tools/predict.ts b/packages/api/src/tools/predict.ts new file mode 100644 index 00000000000..42cc0b32037 --- /dev/null +++ b/packages/api/src/tools/predict.ts @@ -0,0 +1,401 @@ +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'; + +/** 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); +} + +/** Word-boundary matched, so a short name cannot match inside another word. */ +export function namedInRequest(candidates: readonly PredictCandidate[], request: string): string[] { + const haystack = request.toLowerCase(); + if (haystack.length === 0) { + return []; + } + const found: string[] = []; + for (const candidate of candidates) { + const base = baseToolName(candidate.name).toLowerCase(); + if (base.length < 4) { + continue; + } + const at = haystack.indexOf(base); + if (at === -1) { + continue; + } + const before = at === 0 ? '' : haystack[at - 1]; + const after = haystack[at + base.length] ?? ''; + if (/[a-z0-9]/.test(before) || /[a-z0-9]/.test(after)) { + continue; + } + found.push(candidate.name); + } + 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) : []; + + 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, + 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; +} + +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 []; + } + + /** 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, params.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, + }); + + 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 9d7645289a7..04142c1471b 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2937,6 +2937,48 @@ export const classificationSchema = z.object({ /** 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), + 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 carry nothing durable. */ + memoryGate: z + .object({ + enabled: z.boolean().default(false), + threshold: z.number().min(0).max(1).default(0.25), + /** Replaces the durability question. Unset uses the built-in wording. */ + 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(), + }) + .default({}), }); export type TClassificationConfig = z.infer; From 535c89d0e1be5f665c3623f39ad726b8094f86c1 Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Tue, 22 Sep 2026 13:34:57 +0900 Subject: [PATCH 03/11] feat: categorize memories in the gate's existing request --- librechat.example.yaml | 5 + packages/api/src/agents/memory.ts | 23 ++-- packages/api/src/memory/gate.spec.ts | 159 ++++++++++++++++++++++++--- packages/api/src/memory/gate.ts | 101 ++++++++++++++--- packages/data-provider/src/config.ts | 4 +- 5 files changed, 254 insertions(+), 38 deletions(-) diff --git a/librechat.example.yaml b/librechat.example.yaml index e9ed36980a0..2fd8a8bd081 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -571,6 +571,11 @@ actions: # # as booleans. # 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: diff --git a/packages/api/src/agents/memory.ts b/packages/api/src/agents/memory.ts index 4d5c39a48b8..c9258434f85 100644 --- a/packages/api/src/agents/memory.ts +++ b/packages/api/src/agents/memory.ts @@ -1078,13 +1078,20 @@ export async function createMemoryProcessor({ messages: BaseMessage[], inspectionMessages?: BaseMessage[], ): Promise<(TAttachment | null)[] | undefined> { - if (gate != null && !(await gate(messages))) { - logger.debug('[MemoryAgent] Turn carries nothing durable; skipping', { - userId, - conversationId, - messageId, - }); - return undefined; + let turnInstructions = finalInstructions; + if (gate != null) { + const judgment = await gate({ messages, validKeys }); + if (!judgment.process) { + logger.debug('[MemoryAgent] Turn carries nothing durable; skipping', { + userId, + conversationId, + messageId, + }); + return undefined; + } + if (judgment.hint != null) { + turnInstructions = `${finalInstructions}\n\n${judgment.hint}`; + } } try { return await processMemory({ @@ -1105,7 +1112,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/memory/gate.spec.ts b/packages/api/src/memory/gate.spec.ts index bdbc977b8bd..a76a4f3d384 100644 --- a/packages/api/src/memory/gate.spec.ts +++ b/packages/api/src/memory/gate.spec.ts @@ -7,7 +7,13 @@ import type { import type { MemoryGateSettings } from './gate'; import { createMemoryGate, transcribeTail, DURABLE_QUESTION } from './gate'; -const ON: MemoryGateSettings = { enabled: true, threshold: 0.25 }; +const ON: MemoryGateSettings = { + enabled: true, + threshold: 0.25, + categorize: false, + categoryThreshold: 0.4, + detectUpdates: false, +}; function stubClassifier(probability: number | Error): { classifier: Classifier; @@ -82,16 +88,14 @@ 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: { enabled: false, threshold: 0.25 } }), - ).toBeNull(); + expect(createMemoryGate({ classifier, settings: { ...ON, enabled: false } })).toBeNull(); }); it('processes a turn that carries something durable', async () => { const { classifier, requests } = stubClassifier(0.88); const gate = createMemoryGate({ classifier, settings: ON }); - await expect(gate?.(TURN)).resolves.toBe(true); + await expect(gate?.({ messages: TURN })).resolves.toMatchObject({ process: true }); expect(requests).toHaveLength(1); expect(requests[0].questions.durable).toBeDefined(); }); @@ -100,28 +104,30 @@ describe('createMemoryGate', () => { const { classifier } = stubClassifier(0.03); const gate = createMemoryGate({ classifier, settings: ON }); - await expect(gate?.([new HumanMessage('thanks!')])).resolves.toBe(false); + 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?.(TURN)).resolves.toBe(true); + 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: { enabled: true, threshold: 0.8 } }); + const gate = createMemoryGate({ classifier, settings: { ...ON, threshold: 0.8 } }); - await expect(gate?.(TURN)).resolves.toBe(false); + 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?.(TURN)).resolves.toBe(true); + await expect(gate?.({ messages: TURN })).resolves.toMatchObject({ process: true }); }); it('processes the turn when the answer comes back the wrong shape', async () => { @@ -136,14 +142,14 @@ describe('createMemoryGate', () => { }; const gate = createMemoryGate({ classifier, settings: ON }); - await expect(gate?.(TURN)).resolves.toBe(true); + 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?.([])).resolves.toBe(false); + await expect(gate?.({ messages: [] })).resolves.toMatchObject({ process: false }); expect(requests).toHaveLength(0); }); }); @@ -153,7 +159,7 @@ describe('createMemoryGate prompt overrides', () => { const { classifier, requests } = stubClassifier(0.9); const gate = createMemoryGate({ classifier, settings: ON }); - await gate?.(TURN); + await gate?.({ messages: TURN }); const question = requests[0].questions.durable as unknown as { instructions: string; @@ -176,7 +182,7 @@ describe('createMemoryGate prompt overrides', () => { }, }); - await gate?.(TURN); + await gate?.({ messages: TURN }); const question = requests[0].questions.durable as unknown as { instructions: string; @@ -187,3 +193,128 @@ describe('createMemoryGate prompt overrides', () => { 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 durable = { type: 'boolean' as const, probability: 0.9 }; + + it('asks only the durability question by default', async () => { + const { classifier, requests } = stubAnswers({ durable }); + const gate = createMemoryGate({ classifier, settings: ON }); + + await gate?.({ messages: TURN, validKeys: KEYS }); + + expect(Object.keys(requests[0].questions)).toEqual(['durable']); + }); + + it('adds the category and update questions to the same request', async () => { + const { classifier, requests } = stubAnswers({ durable }); + 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', 'durable', 'updates']); + }); + + it('skips categorization when no valid keys are configured', async () => { + const { classifier, requests } = stubAnswers({ durable }); + const gate = createMemoryGate({ classifier, settings: { ...ON, categorize: true } }); + + await gate?.({ messages: TURN }); + + expect(Object.keys(requests[0].questions)).toEqual(['durable']); + }); + + it('suggests the winning key', async () => { + const { classifier } = stubAnswers({ + durable, + 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({ + durable, + 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({ + durable, + 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 is not durable', async () => { + const { classifier } = stubAnswers({ + durable: { 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 index 0246af4f728..809eb21c0aa 100644 --- a/packages/api/src/memory/gate.ts +++ b/packages/api/src/memory/gate.ts @@ -1,12 +1,22 @@ import { logger } from '@librechat/data-schemas'; import type { BaseMessage } from '@librechat/agents/langchain/messages'; import type { TClassificationConfig } from 'librechat-data-provider'; -import type { BooleanQuestion } from '~/classification/types'; -import type { Classifier } from '~/classification/types'; -import { boolean, isTrue } from '~/classification/questions'; -import { isBooleanAnswer } from '~/classification/types'; +import type { Classifier, BooleanQuestion, ClassificationQuestion } from '~/classification/types'; +import { isBooleanAnswer, isChoiceAnswer } from '~/classification/types'; +import { boolean, choice, isTrue } from '~/classification/questions'; -export type MemoryGate = (messages: BaseMessage[]) => Promise; +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']; @@ -20,6 +30,7 @@ export interface CreateMemoryGateParams { const DEFAULT_WINDOW = 6; const DEFAULT_MAX_CHARS = 8_000; +const PROCESS = { process: true } as const; export const DURABLE_QUESTION: BooleanQuestion = boolean( 'Does `conversation` hold something about this user that would still matter in an unrelated ' + @@ -32,6 +43,18 @@ export const DURABLE_QUESTION: BooleanQuestion = boolean( }, ); +export const CATEGORY_INSTRUCTIONS = + 'Which stored memory does the durable part of `conversation` belong under?'; + +export const UPDATE_QUESTION: BooleanQuestion = boolean( + 'Does `conversation` 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') { @@ -85,6 +108,21 @@ export function transcribeTail( return lines.reverse().join('\n'); } +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' + + `The durable part of this turn 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 } = params; if (settings?.enabled !== true) { @@ -93,39 +131,72 @@ export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | n const windowSize = params.windowSize ?? DEFAULT_WINDOW; const maxChars = params.maxChars ?? DEFAULT_MAX_CHARS; const threshold = settings.threshold; - const question = boolean(settings.instructions ?? DURABLE_QUESTION.instructions, { + const durable = boolean(settings.instructions ?? DURABLE_QUESTION.instructions, { true: settings.whenTrue ?? DURABLE_QUESTION.criteria?.true, false: settings.whenFalse ?? DURABLE_QUESTION.criteria?.false, }); - return async function memoryGate(messages: BaseMessage[]): Promise { + function buildQuestions(validKeys: string[]): Record { + const questions: Record = { durable }; + 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 transcript = transcribeTail(messages ?? [], windowSize, maxChars); if (transcript.length === 0) { - return false; + return { process: false }; } try { const response = await classifier.classify({ label: 'memory-gate', signal, state: { conversation: transcript }, - questions: { durable: question }, + questions: buildQuestions(validKeys), }); + const answer = response.answers.durable; if (!isBooleanAnswer(answer)) { - return true; + return PROCESS; + } + if (!isTrue(answer, threshold)) { + logger.debug( + `[memoryGate] durable ${answer.probability.toFixed(2)} below ${threshold}: skipping`, + ); + return { process: false }; } - const keep = isTrue(answer, threshold); + + 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] durable ${answer.probability.toFixed(2)} vs threshold ${threshold}: ` + - `${keep ? 'processing' : 'skipping'} memory`, + `[memoryGate] durable ${answer.probability.toFixed(2)}: processing` + + (key != null ? `, suggesting \`${key}\`${updates ? ' as an update' : ''}` : ''), ); - return keep; + 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 true; + return PROCESS; } }; } diff --git a/packages/data-provider/src/config.ts b/packages/data-provider/src/config.ts index 04142c1471b..d372df77ec8 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2969,7 +2969,6 @@ export const classificationSchema = z.object({ .object({ enabled: z.boolean().default(false), threshold: z.number().min(0).max(1).default(0.25), - /** Replaces the durability question. Unset uses the built-in wording. */ instructions: z.string().min(1).max(4_000).optional(), /** * What a yes and a no mean. Named `whenTrue`/`whenFalse` because YAML @@ -2977,6 +2976,9 @@ export const classificationSchema = z.object({ */ 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({}), }); From c3e2e918be7f532c8137d5b177171ec2cf4d450c Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Tue, 22 Sep 2026 21:27:25 +0900 Subject: [PATCH 04/11] feat: let each capability set its own classifier timeout --- librechat.example.yaml | 3 +++ packages/api/src/memory/gate.ts | 1 + packages/api/src/tools/predict.ts | 1 + packages/data-provider/src/config.ts | 4 ++++ 4 files changed, 9 insertions(+) diff --git a/librechat.example.yaml b/librechat.example.yaml index 2fd8a8bd081..5553f62f853 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -554,6 +554,9 @@ actions: # toolSelection: # enabled: true # shortlist: 5 +# # Overrides the provider timeout. Ranking a large catalog is a much bigger +# # request than a yes/no question. +# timeoutMs: 20000 # # Below this probability that any tool is needed, only tools the request # # names outright are surfaced. # needsToolThreshold: 0.15 diff --git a/packages/api/src/memory/gate.ts b/packages/api/src/memory/gate.ts index 809eb21c0aa..3244e1264fe 100644 --- a/packages/api/src/memory/gate.ts +++ b/packages/api/src/memory/gate.ts @@ -162,6 +162,7 @@ export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | n const response = await classifier.classify({ label: 'memory-gate', signal, + timeoutMs: settings.timeoutMs, state: { conversation: transcript }, questions: buildQuestions(validKeys), }); diff --git a/packages/api/src/tools/predict.ts b/packages/api/src/tools/predict.ts index 42cc0b32037..b923608aa85 100644 --- a/packages/api/src/tools/predict.ts +++ b/packages/api/src/tools/predict.ts @@ -192,6 +192,7 @@ export async function predictTools(params: PredictToolsParams): Promise Date: Thu, 24 Sep 2026 12:14:46 +0900 Subject: [PATCH 05/11] fix: bound each judgment by one deadline and trim unused question types --- librechat.example.yaml | 7 +- .../api/src/classification/config.spec.ts | 9 +-- .../api/src/classification/presets.spec.ts | 21 +++--- .../classification/providers/dialect.spec.ts | 6 +- .../src/classification/providers/dialect.ts | 7 +- .../src/classification/providers/http.spec.ts | 51 +++++++++++++- .../api/src/classification/providers/http.ts | 8 ++- .../src/classification/providers/transport.ts | 37 ++++++++-- .../api/src/classification/questions.spec.ts | 69 +------------------ packages/api/src/classification/questions.ts | 44 ------------ packages/api/src/classification/registry.ts | 15 ++-- packages/api/src/classification/resolve.ts | 2 +- packages/api/src/classification/types.ts | 22 +----- packages/data-provider/src/config.ts | 5 +- 14 files changed, 121 insertions(+), 182 deletions(-) diff --git a/librechat.example.yaml b/librechat.example.yaml index 5553f62f853..42c87f2479f 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -515,9 +515,10 @@ actions: # - 'host.docker.internal:8080' # - '127.0.0.1:8080' -# Classification: small typed judgments (yes/no, pick one, rate) that code can -# branch on. Off unless enabled. The key is read from the environment variable -# named by apiKeyEnv, never written here. +# 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. # classification: # enabled: true # provider: http diff --git a/packages/api/src/classification/config.spec.ts b/packages/api/src/classification/config.spec.ts index b1421305e79..6b6973ebce2 100644 --- a/packages/api/src/classification/config.spec.ts +++ b/packages/api/src/classification/config.spec.ts @@ -1,5 +1,6 @@ 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'; @@ -13,8 +14,8 @@ const ANSWER = { model: 'jev-1.13.0', answers: { d: { type: 'noul', noul: 0.8 } function recorder(response: unknown = ANSWER) { const calls: { url: string; body: Record }[] = []; - const fetch = async (url: string, init: { body?: string }) => { - calls.push({ url, body: JSON.parse(init.body ?? '{}') }); + const fetch: ProviderFetch = async (url, init) => { + calls.push({ url, body: JSON.parse(init.body) }); return { ok: true, status: 200, @@ -22,11 +23,11 @@ function recorder(response: unknown = ANSWER) { text: async () => JSON.stringify(response), }; }; - return { calls, fetch: fetch as never }; + return { calls, fetch }; } function parse(raw: unknown): TClassificationConfig { - return classificationSchema.parse(raw) as TClassificationConfig; + return classificationSchema.parse(raw); } describe('classification config', () => { diff --git a/packages/api/src/classification/presets.spec.ts b/packages/api/src/classification/presets.spec.ts index 0bfcd642ad4..ba2869ca0c8 100644 --- a/packages/api/src/classification/presets.spec.ts +++ b/packages/api/src/classification/presets.spec.ts @@ -1,4 +1,6 @@ -import type { TClassificationConfig } from 'librechat-data-provider'; +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'; @@ -7,8 +9,8 @@ type Captured = { url: string; body: Record }; function recorder(response: unknown) { const calls: Captured[] = []; - const fetch = async (url: string, init: { body?: string }) => { - calls.push({ url, body: JSON.parse(init.body ?? '{}') }); + const fetch: ProviderFetch = async (url, init) => { + calls.push({ url, body: JSON.parse(init.body) }); return { ok: true, status: 200, @@ -16,17 +18,18 @@ function recorder(response: unknown) { text: async () => JSON.stringify(response), }; }; - return { calls, fetch: fetch as never }; + return { calls, fetch }; } -function configFor(provider: string, settings?: Record): TClassificationConfig { - return { +function configFor( + provider: string, + settings?: TClassificationProviderConfig, +): TClassificationConfig { + return classificationSchema.parse({ enabled: true, provider, providers: settings == null ? {} : { [provider]: settings }, - toolSelection: {}, - memoryGate: {}, - } as unknown as TClassificationConfig; + }); } const NOUL = { model: 'jev-1.13.0', answers: { d: { type: 'noul', noul: 0.8 } }, usage: {} }; diff --git a/packages/api/src/classification/providers/dialect.spec.ts b/packages/api/src/classification/providers/dialect.spec.ts index eaeb4c937ab..83cf4d1563d 100644 --- a/packages/api/src/classification/providers/dialect.spec.ts +++ b/packages/api/src/classification/providers/dialect.spec.ts @@ -1,5 +1,5 @@ import { toWireQuestion, readAnswer } from './dialect'; -import { boolean, choice, score } from '../questions'; +import { boolean, choice } from '../questions'; describe('toWireQuestion', () => { it('keeps the port vocabulary by default', () => { @@ -10,12 +10,10 @@ describe('toWireQuestion', () => { expect(toWireQuestion(boolean('durable?'), 'systemone').type).toBe('noul'); }); - it('leaves choice and score alone in either dialect', () => { + it('leaves a choice alone in either dialect', () => { const pick = choice('which', { a: null, b: null }); - const rate = score('how much', ['low', 'high']); expect(toWireQuestion(pick, 'systemone').type).toBe('choice'); - expect(toWireQuestion(rate, 'systemone').type).toBe('score'); expect(toWireQuestion(pick, 'port').criteria).toEqual({ a: null, b: null }); }); diff --git a/packages/api/src/classification/providers/dialect.ts b/packages/api/src/classification/providers/dialect.ts index 1f267a290b0..bcea38561f8 100644 --- a/packages/api/src/classification/providers/dialect.ts +++ b/packages/api/src/classification/providers/dialect.ts @@ -1,6 +1,7 @@ +import type { TClassificationProviderConfig } from 'librechat-data-provider'; import type { ClassificationAnswer, ClassificationQuestion } from '../types'; -export type Dialect = 'port' | 'systemone'; +export type Dialect = NonNullable; interface WireQuestion { type: string; @@ -13,7 +14,6 @@ interface WireAnswer { probability?: unknown; noul?: unknown; choice?: unknown; - score?: unknown; confidence?: unknown; probabilities?: unknown; } @@ -42,8 +42,5 @@ export function readAnswer(answer: unknown, dialect: Dialect): ClassificationAns if (record.type === 'choice' && typeof record.choice === 'string') { return { type: 'choice', choice: record.choice, confidence, probabilities }; } - if (record.type === 'score' && typeof record.score === 'number') { - return { type: 'score', score: record.score, confidence, probabilities }; - } return null; } diff --git a/packages/api/src/classification/providers/http.spec.ts b/packages/api/src/classification/providers/http.spec.ts index 6234637ed4e..79ce8deb494 100644 --- a/packages/api/src/classification/providers/http.spec.ts +++ b/packages/api/src/classification/providers/http.spec.ts @@ -84,7 +84,7 @@ describe('createHttpClassifier', () => { expect(classifier.id).toBe('http'); }); - it('carries choice and score answers through', async () => { + it('carries a choice answer through and drops a type it does not know', async () => { const body = JSON.stringify({ model: 'test-1', answers: { @@ -99,12 +99,11 @@ describe('createHttpClassifier', () => { state: 'x', questions: { pick: { type: 'choice', instructions: 'which', criteria: { a: null, b: null } }, - rate: { type: 'score', instructions: 'how much', criteria: ['low', 'high'] }, }, }); expect(result.answers.pick).toMatchObject({ type: 'choice', choice: 'b', confidence: 0.7 }); - expect(result.answers.rate).toMatchObject({ type: 'score', score: 1.4 }); + expect(result.answers).not.toHaveProperty('rate'); }); it('retries a 429 and honors retry-after', async () => { @@ -138,6 +137,7 @@ describe('createHttpClassifier', () => { apiKey: 'sk-test', endpoint: ENDPOINT, fetch: transport, + timeoutMs: 30_000, sleep: async (ms) => { waits.push(ms); }, @@ -148,6 +148,51 @@ describe('createHttpClassifier', () => { 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, diff --git a/packages/api/src/classification/providers/http.ts b/packages/api/src/classification/providers/http.ts index 407ad65adad..465dbe16bc0 100644 --- a/packages/api/src/classification/providers/http.ts +++ b/packages/api/src/classification/providers/http.ts @@ -14,6 +14,7 @@ import { createTransport } from './transport'; export const PROVIDER_ID = 'http'; export interface HttpProviderOptions { + providerId?: string; apiKey: string; /** Full URL, not a base path. */ endpoint: string; @@ -82,13 +83,14 @@ export function parseEnvelope( } export function createHttpClassifier(options: HttpProviderOptions): Classifier { - const send = createTransport({ providerId: PROVIDER_ID, ...options }); + 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: PROVIDER_ID, + id: providerId, model, async classify(request: ClassificationRequest): Promise { const questions: Record = {}; @@ -106,7 +108,7 @@ export function createHttpClassifier(options: HttpProviderOptions): Classifier { request.label ?? 'classify', request.timeoutMs, ); - return parseEnvelope(body, PROVIDER_ID, (a) => readAnswer(a, dialect), responseKey); + 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 index c18a5302cca..476b96b2b01 100644 --- a/packages/api/src/classification/providers/transport.ts +++ b/packages/api/src/classification/providers/transport.ts @@ -29,7 +29,7 @@ export interface TransportOptions { timeoutMs?: number; maxRetries?: number; fetch?: ProviderFetch; - sleep?: (ms: number) => Promise; + sleep?: (ms: number, signal?: AbortSignal) => Promise; } export type Transport = ( @@ -76,8 +76,20 @@ function briefly(body: string): string { return flat.length > 200 ? `${flat.slice(0, 200)}…` : flat; } -const defaultSleep = (ms: number): Promise => - new Promise((resolve) => setTimeout(resolve, ms)); +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; @@ -161,14 +173,16 @@ export function createTransport(options: TransportOptions): Transport { } } + /** `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, timeoutMs); + 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) { @@ -179,9 +193,18 @@ export function createTransport(options: TransportOptions): Transport { if (attemptNo === maxRetries || !isRetryable(lastError.failure)) { break; } - await sleep( - lastError.retryAfterMs ?? BACKOFF_MS[Math.min(attemptNo, BACKOFF_MS.length - 1)], - ); + 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 ( diff --git a/packages/api/src/classification/questions.spec.ts b/packages/api/src/classification/questions.spec.ts index 7551b1346cc..306c3611d7d 100644 --- a/packages/api/src/classification/questions.spec.ts +++ b/packages/api/src/classification/questions.spec.ts @@ -1,15 +1,5 @@ -import type { ScoreAnswer, ChoiceAnswer, BooleanAnswer } from './types'; -import { - score, - level, - label, - choice, - isTrue, - ranked, - boolean, - normalized, - probabilityOf, -} from './questions'; +import type { ChoiceAnswer, BooleanAnswer } from './types'; +import { choice, isTrue, ranked, boolean } from './questions'; describe('question builders', () => { it('builds a boolean without criteria', () => { @@ -27,17 +17,12 @@ describe('question builders', () => { }); }); - it('builds a choice and a score', () => { + it('builds a choice', () => { expect(choice('Which team?', { billing: 'money', tech: 'bugs' })).toEqual({ type: 'choice', instructions: 'Which team?', criteria: { billing: 'money', tech: 'bugs' }, }); - expect(score('How angry?', ['Calm', 'Cross', 'Furious'])).toEqual({ - type: 'score', - instructions: 'How angry?', - criteria: ['Calm', 'Cross', 'Furious'], - }); }); }); @@ -62,14 +47,6 @@ describe('choice helpers', () => { probabilities: { tech: 0.7, billing: 0.2, sales: 0.05 }, }; - it('reads one option probability', () => { - expect(probabilityOf(answer, 'billing')).toBeCloseTo(0.2); - }); - - it('returns zero for an option that was never offered', () => { - expect(probabilityOf(answer, 'absent')).toBe(0); - }); - it('ranks options most probable first', () => { expect(ranked(answer)).toEqual(['tech', 'billing', 'sales']); }); @@ -82,43 +59,3 @@ describe('choice helpers', () => { expect(ranked(undefined)).toEqual([]); }); }); - -describe('score helpers', () => { - const levels = ['Calm', 'Frustrated', 'Very angry']; - const answer: ScoreAnswer = { - type: 'score', - score: 1.24, - confidence: 0.6, - probabilities: { '0': 0.12, '1': 0.52, '2': 0.36 }, - }; - - it('reports the most probable level, not the rounded score', () => { - expect(level(answer)).toBe(1); - }); - - it('labels that level', () => { - expect(label(answer, levels)).toBe('Frustrated'); - }); - - it('returns an empty label when the levels do not cover it', () => { - expect(label(answer, ['only one'])).toBe(''); - }); - - it('normalizes the weighted score against the top level', () => { - expect(normalized(answer, levels.length)).toBeCloseTo(0.62); - }); - - it('clamps a score outside the level range', () => { - expect(normalized({ ...answer, score: 99 }, levels.length)).toBe(1); - expect(normalized({ ...answer, score: -1 }, levels.length)).toBe(0); - }); - - it('returns zero when there are too few levels to normalize', () => { - expect(normalized(answer, 1)).toBe(0); - }); - - it('handles a missing answer', () => { - expect(level(undefined)).toBe(0); - expect(normalized(undefined, 3)).toBe(0); - }); -}); diff --git a/packages/api/src/classification/questions.ts b/packages/api/src/classification/questions.ts index 471ffe836b8..86e51bf842c 100644 --- a/packages/api/src/classification/questions.ts +++ b/packages/api/src/classification/questions.ts @@ -1,7 +1,5 @@ import type { - ScoreAnswer, ChoiceAnswer, - ScoreQuestion, BooleanAnswer, ChoiceQuestion, BooleanQuestion, @@ -24,24 +22,10 @@ export function choice( return { type: 'choice', instructions, criteria }; } -export function score( - instructions: ClassificationText, - levels: ClassificationText[], -): ScoreQuestion { - return { type: 'score', instructions, criteria: levels }; -} - export function isTrue(answer: BooleanAnswer | undefined, threshold: number): boolean { return answer != null && answer.probability >= threshold; } -export function probabilityOf( - answer: ChoiceAnswer | ScoreAnswer | undefined, - option: string, -): number { - return answer?.probabilities?.[option] ?? 0; -} - /** Options above `floor`, most probable first. */ export function ranked(answer: ChoiceAnswer | undefined, floor = 0): string[] { if (answer == null) { @@ -52,31 +36,3 @@ export function ranked(answer: ChoiceAnswer | undefined, floor = 0): string[] { .sort((a, b) => b[1] - a[1]) .map(([option]) => option); } - -/** The most probable level, which is not always the rounded weighted score. */ -export function level(answer: ScoreAnswer | undefined): number { - if (answer == null) { - return 0; - } - let best = 0; - let bestProbability = -1; - for (const [key, probability] of Object.entries(answer.probabilities)) { - if (probability > bestProbability) { - bestProbability = probability; - best = Number(key); - } - } - return Number.isFinite(best) ? best : 0; -} - -export function label(answer: ScoreAnswer | undefined, levels: readonly string[]): string { - return levels[level(answer)] ?? ''; -} - -/** The weighted score as a fraction of the highest level. */ -export function normalized(answer: ScoreAnswer | undefined, levelCount: number): number { - if (answer == null || levelCount < 2) { - return 0; - } - return Math.min(Math.max(answer.score / (levelCount - 1), 0), 1); -} diff --git a/packages/api/src/classification/registry.ts b/packages/api/src/classification/registry.ts index 65e2ec6a5ae..16ba40de231 100644 --- a/packages/api/src/classification/registry.ts +++ b/packages/api/src/classification/registry.ts @@ -1,18 +1,9 @@ +import type { TClassificationProviderConfig } from 'librechat-data-provider'; import type { ProviderFetch } from './providers/transport'; -import type { Dialect } from './providers/dialect'; import type { Classifier } from './types'; import { createHttpClassifier } from './providers/http'; -export interface ProviderSettings { - baseURL?: string; - model?: string; - dialect?: Dialect; - requestKey?: string; - responseKey?: string; - timeoutMs?: number; - maxRetries?: number; - apiKeyEnv?: string; -} +export type ProviderSettings = TClassificationProviderConfig; export const DEFAULT_API_KEY_ENV = 'CLASSIFIER_API_KEY'; @@ -72,8 +63,10 @@ export function createClassifier( settings: ProviderSettings, apiKey: string, fetch?: ProviderFetch, + providerId?: string, ): Classifier { return createHttpClassifier({ + providerId, apiKey, endpoint: settings.baseURL ?? '', model: settings.model, diff --git a/packages/api/src/classification/resolve.ts b/packages/api/src/classification/resolve.ts index 3964bedf41c..d3a04cc4961 100644 --- a/packages/api/src/classification/resolve.ts +++ b/packages/api/src/classification/resolve.ts @@ -76,7 +76,7 @@ export function resolveClassifier(params: ResolveClassifierParams): Classifier | } try { - const classifier = createClassifier(settings, apiKey, params.fetch); + const classifier = createClassifier(settings, apiKey, params.fetch, providerId); cache.set(config, { apiKey, provider: providerId, classifier }); return classifier; } catch (error) { diff --git a/packages/api/src/classification/types.ts b/packages/api/src/classification/types.ts index 9a3cc9ae7da..a616ca60de2 100644 --- a/packages/api/src/classification/types.ts +++ b/packages/api/src/classification/types.ts @@ -28,13 +28,7 @@ export interface ChoiceQuestion { criteria: Record; } -export interface ScoreQuestion { - type: 'score'; - instructions: ClassificationText; - criteria: ClassificationText[]; -} - -export type ClassificationQuestion = BooleanQuestion | ChoiceQuestion | ScoreQuestion; +export type ClassificationQuestion = BooleanQuestion | ChoiceQuestion; export interface BooleanAnswer { type: 'boolean'; @@ -49,15 +43,7 @@ export interface ChoiceAnswer { probabilities: Record; } -export interface ScoreAnswer { - type: 'score'; - /** 0 to levels - 1. */ - score: number; - confidence: number | null; - probabilities: Record; -} - -export type ClassificationAnswer = BooleanAnswer | ChoiceAnswer | ScoreAnswer; +export type ClassificationAnswer = BooleanAnswer | ChoiceAnswer; export interface ClassificationUsage { inputTokens: number; @@ -122,7 +108,3 @@ export function isBooleanAnswer(answer: ClassificationAnswer | undefined): answe export function isChoiceAnswer(answer: ClassificationAnswer | undefined): answer is ChoiceAnswer { return answer?.type === 'choice'; } - -export function isScoreAnswer(answer: ClassificationAnswer | undefined): answer is ScoreAnswer { - return answer?.type === 'score'; -} diff --git a/packages/data-provider/src/config.ts b/packages/data-provider/src/config.ts index 50b03c4ca28..b556c81056c 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2918,9 +2918,10 @@ export const classificationProviderSchema = z requestKey: z.string().optional(), /** Reads the answer envelope from this key, for hosts that wrap the response. */ responseKey: z.string().optional(), - /** Per-request ceiling. A judgment that misses it is abandoned, never awaited. */ + /** 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 for a rate limit or a server error only. */ + /** 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(), From 4f84f3e68b0cb298e95e9d52d0f3939301da66be Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Thu, 24 Sep 2026 12:17:09 +0900 Subject: [PATCH 06/11] fix: gate memory on the newest turn and apply classification defaults at runtime --- api/server/controllers/agents/client.js | 23 ++---- librechat.example.yaml | 10 ++- packages/api/src/agents/memory.spec.ts | 73 +++++++++++++++++++ packages/api/src/agents/memory.ts | 3 +- packages/api/src/agents/run.ts | 2 +- packages/api/src/memory/index.ts | 1 + packages/api/src/memory/window.spec.ts | 50 +++++++++++++ packages/api/src/memory/window.ts | 20 +++++ packages/api/src/tools/predict.spec.ts | 12 +++ packages/api/src/tools/predict.ts | 47 ++++++++---- packages/data-schemas/src/app/service.spec.ts | 31 +++++++- packages/data-schemas/src/app/service.ts | 21 +++++- packages/data-schemas/src/types/app.ts | 3 +- 13 files changed, 256 insertions(+), 40 deletions(-) create mode 100644 packages/api/src/memory/window.spec.ts create mode 100644 packages/api/src/memory/window.ts diff --git a/api/server/controllers/agents/client.js b/api/server/controllers/agents/client.js index e00edc016e1..52531a25745 100644 --- a/api/server/controllers/agents/client.js +++ b/api/server/controllers/agents/client.js @@ -47,6 +47,7 @@ const { getRunDiscoveredTools, predictToolsForTurn, createMemoryGate, + selectMemoryWindow, classificationCapability, captureResumeModelParameters, pickResumeContext, @@ -1855,8 +1856,10 @@ class AgentClient extends BaseClient { return wiring; } - /** Builds the independently opt-in live reasoning-label controller. */ - /** @returns {import('@librechat/api').MemoryGate | null} */ + /** + * 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'); @@ -1869,6 +1872,7 @@ class AgentClient extends BaseClient { }); } + /** Builds the independently opt-in live reasoning-label controller. */ buildReasoningLabelWiring(streamId, abortSignal, seedFromContent = false) { if (!streamId || typeof Run?.prototype?.generateReasoningLabel !== 'function') { return undefined; @@ -3525,20 +3529,7 @@ class AgentClient extends BaseClient { */ const chatMessages = messages.filter((m) => !isSkillPrimeMessage(m)); - let messagesToProcess = [...chatMessages]; - if (chatMessages.length > messageWindowSize) { - for (let i = chatMessages.length - messageWindowSize; i >= 0; i--) { - const potentialWindow = chatMessages.slice(i, i + messageWindowSize); - if (potentialWindow[0]?.role === 'user') { - messagesToProcess = [...potentialWindow]; - break; - } - } - - if (messagesToProcess.length === chatMessages.length) { - messagesToProcess = [...chatMessages.slice(-messageWindowSize)]; - } - } + const messagesToProcess = selectMemoryWindow(chatMessages, messageWindowSize); const filteredMessages = messagesToProcess.map((msg) => this.filterImageUrls(msg)); const bufferString = getBufferString(filteredMessages); diff --git a/librechat.example.yaml b/librechat.example.yaml index 42c87f2479f..1fa72ce0908 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -518,7 +518,9 @@ actions: # 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. +# 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 @@ -555,9 +557,9 @@ actions: # toolSelection: # enabled: true # shortlist: 5 -# # Overrides the provider timeout. Ranking a large catalog is a much bigger -# # request than a yes/no question. -# timeoutMs: 20000 +# # 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 diff --git a/packages/api/src/agents/memory.spec.ts b/packages/api/src/agents/memory.spec.ts index 918f9797821..c202454cd0f 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(durable: number) { + const requests: ClassificationRequest[] = []; + const classifier: Classifier = { + id: 'recording', + model: 'test', + async classify(request) { + requests.push(request); + return { + model: 'test', + answers: { durable: { type: 'boolean', probability: durable } }, + 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 holds nothing durable', 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 c9258434f85..d5c6e0375d6 100644 --- a/packages/api/src/agents/memory.ts +++ b/packages/api/src/agents/memory.ts @@ -1080,7 +1080,8 @@ export async function createMemoryProcessor({ ): Promise<(TAttachment | null)[] | undefined> { let turnInstructions = finalInstructions; if (gate != null) { - const judgment = await gate({ messages, validKeys }); + /** `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 carries nothing durable; skipping', { userId, diff --git a/packages/api/src/agents/run.ts b/packages/api/src/agents/run.ts index 34d2dbeea60..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) { diff --git a/packages/api/src/memory/index.ts b/packages/api/src/memory/index.ts index 5c7cf4b3f3d..870b35c0724 100644 --- a/packages/api/src/memory/index.ts +++ b/packages/api/src/memory/index.ts @@ -3,3 +3,4 @@ export * from './authorization'; export * from './handlers'; export * from './protection'; export * from './gate'; +export * from './window'; diff --git a/packages/api/src/memory/window.spec.ts b/packages/api/src/memory/window.spec.ts new file mode 100644 index 00000000000..cd70725705d --- /dev/null +++ b/packages/api/src/memory/window.spec.ts @@ -0,0 +1,50 @@ +import { selectMemoryWindow } from './window'; + +const turn = (role: string, id: string) => ({ role, id }); + +describe('selectMemoryWindow', () => { + it('keeps every message when the conversation fits the window', () => { + const messages = [turn('user', 'u1'), turn('assistant', 'a1')]; + + expect(selectMemoryWindow(messages, 5)).toEqual(messages); + }); + + it('always ends at the newest message after a tool-heavy turn', () => { + const messages = [ + turn('user', 'u1'), + turn('assistant', 'call'), + turn('tool', 't1'), + turn('tool', 't2'), + turn('assistant', 'a1'), + turn('user', 'u2'), + ]; + + expect(selectMemoryWindow(messages, 5).map((m) => m.id)).toEqual(['u2']); + }); + + it('opens the window on the earliest user turn inside it', () => { + const messages = [ + turn('user', 'u1'), + turn('assistant', 'a1'), + turn('user', 'u2'), + turn('assistant', 'a2'), + turn('user', 'u3'), + turn('assistant', 'a3'), + turn('user', 'u4'), + ]; + + expect(selectMemoryWindow(messages, 5).map((m) => m.id)).toEqual([ + 'u2', + 'a2', + 'u3', + 'a3', + 'u4', + ]); + }); + + it('falls back to the plain tail when no user turn is inside it', () => { + const messages = [turn('user', 'u1'), ...['a', 'b', 'c'].map((id) => turn('tool', id))]; + + expect(selectMemoryWindow(messages, 2).map((m) => m.id)).toEqual(['b', 'c']); + }); +}); diff --git a/packages/api/src/memory/window.ts b/packages/api/src/memory/window.ts new file mode 100644 index 00000000000..9b2984f1cd5 --- /dev/null +++ b/packages/api/src/memory/window.ts @@ -0,0 +1,20 @@ +interface RoleTagged { + role?: string; +} + +/** Last `size` messages, trimmed to open on a user turn. Always ends at the newest message. */ +export function selectMemoryWindow( + messages: readonly T[], + size: number, +): T[] { + if (messages.length <= size) { + return [...messages]; + } + const start = messages.length - size; + for (let i = start; i < messages.length; i++) { + if (messages[i]?.role === 'user') { + return messages.slice(i); + } + } + return messages.slice(start); +} diff --git a/packages/api/src/tools/predict.spec.ts b/packages/api/src/tools/predict.spec.ts index f8b8318e2c4..4097f8b11ad 100644 --- a/packages/api/src/tools/predict.spec.ts +++ b/packages/api/src/tools/predict.spec.ts @@ -319,6 +319,18 @@ describe('namedInRequest', () => { 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', () => { diff --git a/packages/api/src/tools/predict.ts b/packages/api/src/tools/predict.ts index b923608aa85..f8bd9960d14 100644 --- a/packages/api/src/tools/predict.ts +++ b/packages/api/src/tools/predict.ts @@ -11,6 +11,7 @@ import type { 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__'; @@ -91,28 +92,40 @@ export function baseToolName(name: string): string { return index === -1 ? name : name.slice(0, index); } -/** Word-boundary matched, so a short name cannot match inside another word. */ -export function namedInRequest(candidates: readonly PredictCandidate[], request: string): string[] { +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) { + 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) { - continue; - } - const at = haystack.indexOf(base); - if (at === -1) { - continue; - } - const before = at === 0 ? '' : haystack[at - 1]; - const after = haystack[at + base.length] ?? ''; - if (/[a-z0-9]/.test(before) || /[a-z0-9]/.test(after)) { + if (base.length < 4 || !mentionsWord(haystack, base)) { continue; } found.push(candidate.name); + if (found.length >= limit) { + break; + } } return found; } @@ -180,7 +193,7 @@ export async function predictTools(params: PredictToolsParams): Promise(); const candidates: PredictCandidate[] = []; for (const agent of params.agents) { - for (const candidate of deferredCandidates(agent.toolRegistry, params.alreadyLoaded)) { + for (const candidate of deferredCandidates(agent.toolRegistry, alreadyLoaded)) { if (seen.has(candidate.name)) { continue; } diff --git a/packages/data-schemas/src/app/service.spec.ts b/packages/data-schemas/src/app/service.spec.ts index 820335f4b5d..ce16fd6b256 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,30 @@ 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); + }); + + 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 021c8763f55..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,7 +180,7 @@ export const AppService = async (params?: { const mcpServersConfig = config.mcpServers || null; const mcpSettings = config.mcpSettings || null; - const classification = config.classification || null; + const classification = loadClassificationConfig(config); const actions = config.actions; const registration = config.registration ?? configDefaults.registration; const interfaceConfig = await loadDefaultInterface({ config, configDefaults }); diff --git a/packages/data-schemas/src/types/app.ts b/packages/data-schemas/src/types/app.ts index fc188edb50e..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 = { @@ -103,7 +104,7 @@ export interface AppConfig { /** MCP settings (domain allowlist, etc.) */ mcpSettings?: TCustomConfig['mcpSettings'] | null; /** Classification provider and the capabilities that consult it */ - classification?: TCustomConfig['classification'] | null; + classification?: TClassificationConfig | null; /** File configuration */ fileConfig?: TFileConfig; /** Secure image links configuration, enabled unless explicitly disabled */ From 83b9c093a0b0e3ead946fbd96e5f147c1636961c Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Thu, 24 Sep 2026 13:00:58 +0900 Subject: [PATCH 07/11] fix: judge only the newest user message in the memory gate --- packages/api/src/memory/gate.spec.ts | 28 +++++++++++++++++++++++ packages/api/src/memory/gate.ts | 34 +++++++++++++++++++++------- 2 files changed, 54 insertions(+), 8 deletions(-) diff --git a/packages/api/src/memory/gate.spec.ts b/packages/api/src/memory/gate.spec.ts index a76a4f3d384..ec2b9fbe40a 100644 --- a/packages/api/src/memory/gate.spec.ts +++ b/packages/api/src/memory/gate.spec.ts @@ -152,6 +152,34 @@ describe('createMemoryGate', () => { 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 prompt overrides', () => { diff --git a/packages/api/src/memory/gate.ts b/packages/api/src/memory/gate.ts index 3244e1264fe..18011c9c4ab 100644 --- a/packages/api/src/memory/gate.ts +++ b/packages/api/src/memory/gate.ts @@ -33,21 +33,22 @@ const DEFAULT_MAX_CHARS = 8_000; const PROCESS = { process: true } as const; export const DURABLE_QUESTION: BooleanQuestion = boolean( - 'Does `conversation` hold something about this user that would still matter in an unrelated ' + - 'conversation weeks from now?', + "Does `latest`, the user's newest message, tell us something about this user that would " + + 'still matter in an unrelated conversation weeks from now? `conversation` is the earlier ' + + 'context and is only there to help read `latest`.', { true: 'A lasting preference, a fact about who they are or what they work on, or a decision ' + 'they want remembered.', - false: 'Small talk, or a detail that only matters inside this task.', + false: 'Small talk, a request for this task, or a detail that only matters inside this task.', }, ); export const CATEGORY_INSTRUCTIONS = - 'Which stored memory does the durable part of `conversation` belong under?'; + 'Which stored memory does the durable part of `latest` belong under?'; export const UPDATE_QUESTION: BooleanQuestion = boolean( - 'Does `conversation` change something already known about this user, rather than adding ' + + '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.', @@ -108,6 +109,16 @@ export function transcribeTail( 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; @@ -154,16 +165,23 @@ export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | n messages, validKeys = [], }: MemoryGateInput): Promise { - const transcript = transcribeTail(messages ?? [], windowSize, maxChars); - if (transcript.length === 0) { + 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: { conversation: transcript }, + state: { latest, conversation: transcript }, questions: buildQuestions(validKeys), }); From 0f3e7d25caa0b39deeccd8e08becbb6632599391 Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Thu, 24 Sep 2026 13:18:06 +0900 Subject: [PATCH 08/11] refactor: move the memory window fix to #16276 --- api/server/controllers/agents/client.js | 16 +++++++- packages/api/src/memory/index.ts | 1 - packages/api/src/memory/window.spec.ts | 50 ------------------------- packages/api/src/memory/window.ts | 20 ---------- 4 files changed, 14 insertions(+), 73 deletions(-) delete mode 100644 packages/api/src/memory/window.spec.ts delete mode 100644 packages/api/src/memory/window.ts diff --git a/api/server/controllers/agents/client.js b/api/server/controllers/agents/client.js index 52531a25745..b75d771fac6 100644 --- a/api/server/controllers/agents/client.js +++ b/api/server/controllers/agents/client.js @@ -47,7 +47,6 @@ const { getRunDiscoveredTools, predictToolsForTurn, createMemoryGate, - selectMemoryWindow, classificationCapability, captureResumeModelParameters, pickResumeContext, @@ -3529,7 +3528,20 @@ class AgentClient extends BaseClient { */ const chatMessages = messages.filter((m) => !isSkillPrimeMessage(m)); - const messagesToProcess = selectMemoryWindow(chatMessages, messageWindowSize); + let messagesToProcess = [...chatMessages]; + if (chatMessages.length > messageWindowSize) { + for (let i = chatMessages.length - messageWindowSize; i >= 0; i--) { + const potentialWindow = chatMessages.slice(i, i + messageWindowSize); + if (potentialWindow[0]?.role === 'user') { + messagesToProcess = [...potentialWindow]; + break; + } + } + + if (messagesToProcess.length === chatMessages.length) { + messagesToProcess = [...chatMessages.slice(-messageWindowSize)]; + } + } const filteredMessages = messagesToProcess.map((msg) => this.filterImageUrls(msg)); const bufferString = getBufferString(filteredMessages); diff --git a/packages/api/src/memory/index.ts b/packages/api/src/memory/index.ts index 870b35c0724..5c7cf4b3f3d 100644 --- a/packages/api/src/memory/index.ts +++ b/packages/api/src/memory/index.ts @@ -3,4 +3,3 @@ export * from './authorization'; export * from './handlers'; export * from './protection'; export * from './gate'; -export * from './window'; diff --git a/packages/api/src/memory/window.spec.ts b/packages/api/src/memory/window.spec.ts deleted file mode 100644 index cd70725705d..00000000000 --- a/packages/api/src/memory/window.spec.ts +++ /dev/null @@ -1,50 +0,0 @@ -import { selectMemoryWindow } from './window'; - -const turn = (role: string, id: string) => ({ role, id }); - -describe('selectMemoryWindow', () => { - it('keeps every message when the conversation fits the window', () => { - const messages = [turn('user', 'u1'), turn('assistant', 'a1')]; - - expect(selectMemoryWindow(messages, 5)).toEqual(messages); - }); - - it('always ends at the newest message after a tool-heavy turn', () => { - const messages = [ - turn('user', 'u1'), - turn('assistant', 'call'), - turn('tool', 't1'), - turn('tool', 't2'), - turn('assistant', 'a1'), - turn('user', 'u2'), - ]; - - expect(selectMemoryWindow(messages, 5).map((m) => m.id)).toEqual(['u2']); - }); - - it('opens the window on the earliest user turn inside it', () => { - const messages = [ - turn('user', 'u1'), - turn('assistant', 'a1'), - turn('user', 'u2'), - turn('assistant', 'a2'), - turn('user', 'u3'), - turn('assistant', 'a3'), - turn('user', 'u4'), - ]; - - expect(selectMemoryWindow(messages, 5).map((m) => m.id)).toEqual([ - 'u2', - 'a2', - 'u3', - 'a3', - 'u4', - ]); - }); - - it('falls back to the plain tail when no user turn is inside it', () => { - const messages = [turn('user', 'u1'), ...['a', 'b', 'c'].map((id) => turn('tool', id))]; - - expect(selectMemoryWindow(messages, 2).map((m) => m.id)).toEqual(['b', 'c']); - }); -}); diff --git a/packages/api/src/memory/window.ts b/packages/api/src/memory/window.ts deleted file mode 100644 index 9b2984f1cd5..00000000000 --- a/packages/api/src/memory/window.ts +++ /dev/null @@ -1,20 +0,0 @@ -interface RoleTagged { - role?: string; -} - -/** Last `size` messages, trimmed to open on a user turn. Always ends at the newest message. */ -export function selectMemoryWindow( - messages: readonly T[], - size: number, -): T[] { - if (messages.length <= size) { - return [...messages]; - } - const start = messages.length - size; - for (let i = start; i < messages.length; i++) { - if (messages[i]?.role === 'user') { - return messages.slice(i); - } - } - return messages.slice(start); -} From 461a47ffb5194b37eb09ef5e31224a6ff37416d9 Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Thu, 24 Sep 2026 13:39:12 +0900 Subject: [PATCH 09/11] fix: gate memory on explicit remember, update and forget requests --- librechat.example.yaml | 11 +++-- packages/api/src/agents/memory.spec.ts | 6 +-- packages/api/src/agents/memory.ts | 2 +- packages/api/src/memory/gate.spec.ts | 46 +++++++++---------- packages/api/src/memory/gate.ts | 36 ++++++++------- packages/data-provider/src/config.ts | 4 +- packages/data-schemas/src/app/service.spec.ts | 1 + 7 files changed, 56 insertions(+), 50 deletions(-) diff --git a/librechat.example.yaml b/librechat.example.yaml index 1fa72ce0908..4d508495160 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -569,12 +569,15 @@ actions: # instructions: Which tool should the assistant call first? # guidance: Prefer the tool whose purpose matches the request. # -# # Skips the memory model on turns that carry nothing durable. +# # 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.25 -# # whenTrue/whenFalse rather than true/false: YAML reads those bare keys -# # as booleans. +# 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 diff --git a/packages/api/src/agents/memory.spec.ts b/packages/api/src/agents/memory.spec.ts index c202454cd0f..c7fda4c0601 100644 --- a/packages/api/src/agents/memory.spec.ts +++ b/packages/api/src/agents/memory.spec.ts @@ -780,7 +780,7 @@ describe('createMemoryTool tokenLimit enforcement', () => { }); describe('memory gate inside the memory processor', () => { - function recordingGate(durable: number) { + function recordingGate(probability: number) { const requests: ClassificationRequest[] = []; const classifier: Classifier = { id: 'recording', @@ -789,7 +789,7 @@ describe('memory gate inside the memory processor', () => { requests.push(request); return { model: 'test', - answers: { durable: { type: 'boolean', probability: durable } }, + answers: { request: { type: 'boolean', probability } }, usage: { inputTokens: 0, outputTokens: 0 }, }; }, @@ -838,7 +838,7 @@ describe('memory gate inside the memory processor', () => { expect(conversation).toContain('I prefer answers in Japanese from now on'); }); - it('skips the memory model when the gate says the turn holds nothing durable', async () => { + 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; diff --git a/packages/api/src/agents/memory.ts b/packages/api/src/agents/memory.ts index d5c6e0375d6..590e309f6cb 100644 --- a/packages/api/src/agents/memory.ts +++ b/packages/api/src/agents/memory.ts @@ -1083,7 +1083,7 @@ export async function createMemoryProcessor({ /** `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 carries nothing durable; skipping', { + logger.debug('[MemoryAgent] Turn asks for no memory change; skipping', { userId, conversationId, messageId, diff --git a/packages/api/src/memory/gate.spec.ts b/packages/api/src/memory/gate.spec.ts index ec2b9fbe40a..1b7268b5e70 100644 --- a/packages/api/src/memory/gate.spec.ts +++ b/packages/api/src/memory/gate.spec.ts @@ -5,7 +5,7 @@ import type { ClassificationRequest, } from '~/classification/types'; import type { MemoryGateSettings } from './gate'; -import { createMemoryGate, transcribeTail, DURABLE_QUESTION } from './gate'; +import { createMemoryGate, transcribeTail, MEMORY_REQUEST_QUESTION } from './gate'; const ON: MemoryGateSettings = { enabled: true, @@ -30,7 +30,7 @@ function stubClassifier(probability: number | Error): { } return { model: 'stub-1', - answers: { durable: { type: 'boolean', probability } }, + answers: { request: { type: 'boolean', probability } }, usage: { inputTokens: 40, outputTokens: 4 }, }; }, @@ -91,16 +91,16 @@ describe('createMemoryGate', () => { expect(createMemoryGate({ classifier, settings: { ...ON, enabled: false } })).toBeNull(); }); - it('processes a turn that carries something durable', async () => { + 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.durable).toBeDefined(); + expect(requests[0].questions.request).toBeDefined(); }); - it('skips a turn that carries nothing durable', async () => { + it('skips a turn that asks for no memory change', async () => { const { classifier } = stubClassifier(0.03); const gate = createMemoryGate({ classifier, settings: ON }); @@ -136,7 +136,7 @@ describe('createMemoryGate', () => { model: 'stub-1', classify: async () => ({ model: 'stub-1', - answers: { durable: { type: 'choice', choice: 'yes', confidence: 1, probabilities: {} } }, + answers: { request: { type: 'choice', choice: 'yes', confidence: 1, probabilities: {} } }, usage: { inputTokens: 1, outputTokens: 1 }, }), }; @@ -189,13 +189,13 @@ describe('createMemoryGate prompt overrides', () => { await gate?.({ messages: TURN }); - const question = requests[0].questions.durable as unknown as { + const question = requests[0].questions.request as unknown as { instructions: string; criteria: { true: string; false: string }; }; - expect(question.instructions).toBe(DURABLE_QUESTION.instructions); - expect(question.criteria.true).toBe(DURABLE_QUESTION.criteria?.true); - expect(question.criteria.false).toBe(DURABLE_QUESTION.criteria?.false); + 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 () => { @@ -212,7 +212,7 @@ describe('createMemoryGate prompt overrides', () => { await gate?.({ messages: TURN }); - const question = requests[0].questions.durable as unknown as { + const question = requests[0].questions.request as unknown as { instructions: string; criteria: { true: string; false: string }; }; @@ -241,19 +241,19 @@ describe('memory classification', () => { return { classifier, requests }; } - const durable = { type: 'boolean' as const, probability: 0.9 }; + const request = { type: 'boolean' as const, probability: 0.9 }; it('asks only the durability question by default', async () => { - const { classifier, requests } = stubAnswers({ durable }); + const { classifier, requests } = stubAnswers({ request }); const gate = createMemoryGate({ classifier, settings: ON }); await gate?.({ messages: TURN, validKeys: KEYS }); - expect(Object.keys(requests[0].questions)).toEqual(['durable']); + expect(Object.keys(requests[0].questions)).toEqual(['request']); }); it('adds the category and update questions to the same request', async () => { - const { classifier, requests } = stubAnswers({ durable }); + const { classifier, requests } = stubAnswers({ request }); const gate = createMemoryGate({ classifier, settings: { ...ON, categorize: true, detectUpdates: true }, @@ -262,21 +262,21 @@ describe('memory classification', () => { await gate?.({ messages: TURN, validKeys: KEYS }); expect(requests).toHaveLength(1); - expect(Object.keys(requests[0].questions).sort()).toEqual(['category', 'durable', 'updates']); + 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({ durable }); + const { classifier, requests } = stubAnswers({ request }); const gate = createMemoryGate({ classifier, settings: { ...ON, categorize: true } }); await gate?.({ messages: TURN }); - expect(Object.keys(requests[0].questions)).toEqual(['durable']); + expect(Object.keys(requests[0].questions)).toEqual(['request']); }); it('suggests the winning key', async () => { const { classifier } = stubAnswers({ - durable, + request, category: { type: 'choice', choice: 'work_context', @@ -294,7 +294,7 @@ describe('memory classification', () => { it('drops a key it is not confident about', async () => { const { classifier } = stubAnswers({ - durable, + request, category: { type: 'choice', choice: 'work_context', @@ -312,7 +312,7 @@ describe('memory classification', () => { it('says so when the turn changes an existing fact', async () => { const { classifier } = stubAnswers({ - durable, + request, category: { type: 'choice', choice: 'preferences', @@ -331,9 +331,9 @@ describe('memory classification', () => { expect(judgment?.hint).toContain('a change to what is already stored'); }); - it('carries no hint when the turn is not durable', async () => { + it('carries no hint when the turn asks for nothing', async () => { const { classifier } = stubAnswers({ - durable: { type: 'boolean', probability: 0.01 }, + request: { type: 'boolean', probability: 0.01 }, category: { type: 'choice', choice: 'personal', diff --git a/packages/api/src/memory/gate.ts b/packages/api/src/memory/gate.ts index 18011c9c4ab..6f5c7cc3541 100644 --- a/packages/api/src/memory/gate.ts +++ b/packages/api/src/memory/gate.ts @@ -32,20 +32,22 @@ const DEFAULT_WINDOW = 6; const DEFAULT_MAX_CHARS = 8_000; const PROCESS = { process: true } as const; -export const DURABLE_QUESTION: BooleanQuestion = boolean( - "Does `latest`, the user's newest message, tell us something about this user that would " + - 'still matter in an unrelated conversation weeks from now? `conversation` is the earlier ' + - 'context and is only there to help read `latest`.', +/** 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: - 'A lasting preference, a fact about who they are or what they work on, or a decision ' + - 'they want remembered.', - false: 'Small talk, a request for this task, or a detail that only matters inside this task.', + '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 durable part of `latest` belong under?'; +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 ' + @@ -128,7 +130,7 @@ export function buildHint(key: string | null, updates: boolean): string | undefi : ''; return ( '\n' + - `The durable part of this turn most likely belongs under \`${key}\`.${change}\n` + + `This request most likely belongs under \`${key}\`.${change}\n` + 'Ignore this if it does not fit what the user actually said.\n' + '' ); @@ -142,13 +144,13 @@ export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | n const windowSize = params.windowSize ?? DEFAULT_WINDOW; const maxChars = params.maxChars ?? DEFAULT_MAX_CHARS; const threshold = settings.threshold; - const durable = boolean(settings.instructions ?? DURABLE_QUESTION.instructions, { - true: settings.whenTrue ?? DURABLE_QUESTION.criteria?.true, - false: settings.whenFalse ?? DURABLE_QUESTION.criteria?.false, + 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 = { durable }; + const questions: Record = { request }; if (settings.categorize === true && validKeys.length > 0) { questions.category = choice( CATEGORY_INSTRUCTIONS, @@ -185,13 +187,13 @@ export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | n questions: buildQuestions(validKeys), }); - const answer = response.answers.durable; + const answer = response.answers.request; if (!isBooleanAnswer(answer)) { return PROCESS; } if (!isTrue(answer, threshold)) { logger.debug( - `[memoryGate] durable ${answer.probability.toFixed(2)} below ${threshold}: skipping`, + `[memoryGate] request ${answer.probability.toFixed(2)} below ${threshold}: skipping`, ); return { process: false }; } @@ -206,7 +208,7 @@ export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | n const updates = isBooleanAnswer(updatesAnswer) && updatesAnswer.probability >= 0.5; logger.debug( - `[memoryGate] durable ${answer.probability.toFixed(2)}: processing` + + `[memoryGate] request ${answer.probability.toFixed(2)}: processing` + (key != null ? `, suggesting \`${key}\`${updates ? ' as an update' : ''}` : ''), ); return { process: true, hint: buildHint(key, updates) }; diff --git a/packages/data-provider/src/config.ts b/packages/data-provider/src/config.ts index b556c81056c..89a4c6244b0 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2968,12 +2968,12 @@ export const classificationSchema = z.object({ needsToolInstructions: z.string().min(1).max(4_000).optional(), }) .default({}), - /** Skips the memory model on turns that carry nothing durable. */ + /** 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.25), + 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 diff --git a/packages/data-schemas/src/app/service.spec.ts b/packages/data-schemas/src/app/service.spec.ts index ce16fd6b256..a5e330f92ee 100644 --- a/packages/data-schemas/src/app/service.spec.ts +++ b/packages/data-schemas/src/app/service.spec.ts @@ -106,6 +106,7 @@ describe('loadClassificationConfig', () => { 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', () => { From 48fc04751b7fc3b6922bd5adda1354084e212ac2 Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Thu, 24 Sep 2026 14:13:26 +0900 Subject: [PATCH 10/11] perf: run tool prediction alongside the memory run --- api/server/controllers/agents/client.js | 19 +++++-- api/server/controllers/agents/client.test.js | 58 ++++++++++++++++++++ 2 files changed, 71 insertions(+), 6 deletions(-) diff --git a/api/server/controllers/agents/client.js b/api/server/controllers/agents/client.js index b75d771fac6..67fd6be4bba 100644 --- a/api/server/controllers/agents/client.js +++ b/api/server/controllers/agents/client.js @@ -4830,17 +4830,19 @@ class AgentClient extends BaseClient { agents.push(...this.agentConfigs.values()); } - /** Ahead of the checkpoint setup below: the prune and `createRun` are - * deliberately overlapped, and an await between them serializes both. */ - const predictedToolNames = await predictToolsForTurn({ + /** Awaited after the memory run starts, so the two overlap. */ + const predictionPromise = predictToolsForTurn({ config: appConfig?.classification, agents, messages, signal: abortController.signal, + }).catch((error) => { + logger.warn( + '[AgentClient] Tool prediction failed; continuing without it', + getSafeErrorMetadata(error), + ); + return []; }); - if (predictedToolNames.length > 0) { - this.predictedToolNames = predictedToolNames; - } const modelBoundCallback = AgentClient.prototype.createModelBoundChatModelCallback.call(this); const initialModelBoundAdmission = @@ -4883,6 +4885,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; diff --git a/api/server/controllers/agents/client.test.js b/api/server/controllers/agents/client.test.js index c646df0436d..5257f205f16 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,60 @@ describe('AgentClient - startup telemetry', () => { }, ); + 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) => From d7e83d78bd3a82fee709e21671a729020c2a1f00 Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Thu, 24 Sep 2026 14:17:53 +0900 Subject: [PATCH 11/11] feat: bill classifier usage as a classification transaction --- api/server/controllers/agents/client.js | 19 ++++++++ api/server/controllers/agents/client.test.js | 35 ++++++++++++++ packages/api/src/memory/gate.spec.ts | 22 +++++++++ packages/api/src/memory/gate.ts | 12 ++++- packages/api/src/tools/predict.spec.ts | 51 ++++++++++++++++++++ packages/api/src/tools/predict.ts | 6 +++ packages/data-schemas/src/methods/tx.spec.ts | 8 +++ packages/data-schemas/src/methods/tx.ts | 2 + 8 files changed, 153 insertions(+), 2 deletions(-) diff --git a/api/server/controllers/agents/client.js b/api/server/controllers/agents/client.js index 67fd6be4bba..2360dd29da3 100644 --- a/api/server/controllers/agents/client.js +++ b/api/server/controllers/agents/client.js @@ -1868,6 +1868,24 @@ class AgentClient extends BaseClient { 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)); }); } @@ -4836,6 +4854,7 @@ class AgentClient extends BaseClient { agents, messages, signal: abortController.signal, + onUsage: (usage, model) => this.recordClassifierUsage(usage, model), }).catch((error) => { logger.warn( '[AgentClient] Tool prediction failed; continuing without it', diff --git a/api/server/controllers/agents/client.test.js b/api/server/controllers/agents/client.test.js index 5257f205f16..a32b867ca32 100644 --- a/api/server/controllers/agents/client.test.js +++ b/api/server/controllers/agents/client.test.js @@ -2381,6 +2381,41 @@ 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( diff --git a/packages/api/src/memory/gate.spec.ts b/packages/api/src/memory/gate.spec.ts index 1b7268b5e70..6d3de34fd75 100644 --- a/packages/api/src/memory/gate.spec.ts +++ b/packages/api/src/memory/gate.spec.ts @@ -182,6 +182,28 @@ describe('createMemoryGate', () => { }); }); +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); diff --git a/packages/api/src/memory/gate.ts b/packages/api/src/memory/gate.ts index 6f5c7cc3541..7d4cdf547c6 100644 --- a/packages/api/src/memory/gate.ts +++ b/packages/api/src/memory/gate.ts @@ -1,7 +1,12 @@ 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, ClassificationQuestion } from '~/classification/types'; +import type { + Classifier, + BooleanQuestion, + ClassificationUsage, + ClassificationQuestion, +} from '~/classification/types'; import { isBooleanAnswer, isChoiceAnswer } from '~/classification/types'; import { boolean, choice, isTrue } from '~/classification/questions'; @@ -26,6 +31,8 @@ export interface CreateMemoryGateParams { 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; @@ -137,7 +144,7 @@ export function buildHint(key: string | null, updates: boolean): string | undefi } export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | null { - const { classifier, settings, signal } = params; + const { classifier, settings, signal, onUsage } = params; if (settings?.enabled !== true) { return null; } @@ -186,6 +193,7 @@ export function createMemoryGate(params: CreateMemoryGateParams): MemoryGate | n state: { latest, conversation: transcript }, questions: buildQuestions(validKeys), }); + onUsage?.(response.usage, classifier.model); const answer = response.answers.request; if (!isBooleanAnswer(answer)) { diff --git a/packages/api/src/tools/predict.spec.ts b/packages/api/src/tools/predict.spec.ts index 4097f8b11ad..d174885f7c8 100644 --- a/packages/api/src/tools/predict.spec.ts +++ b/packages/api/src/tools/predict.spec.ts @@ -1,3 +1,5 @@ +import { classificationSchema } from 'librechat-data-provider'; +import { HumanMessage } from '@librechat/agents/langchain/messages'; import type { LCToolRegistry, LCTool } from '@librechat/agents'; import type { Classifier, @@ -11,6 +13,7 @@ import { namedInRequest, batchCandidates, deferredCandidates, + predictToolsForTurn, RANKING_QUESTION, NEEDS_TOOL_QUESTION, } from './predict'; @@ -478,3 +481,51 @@ describe('shortlistSize with unmeasured confidence', () => { 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 index f8bd9960d14..84167c7d0a7 100644 --- a/packages/api/src/tools/predict.ts +++ b/packages/api/src/tools/predict.ts @@ -361,6 +361,8 @@ export interface PredictToolsForTurnParams { 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 { @@ -410,6 +412,10 @@ export async function predictToolsForTurn(params: PredictToolsForTurnParams): Pr 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(', ')}` : ''), 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 },