From 6233c069397c338abcf909ffdb5546a81745a99d Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Tue, 22 Sep 2026 13:04:21 +0900 Subject: [PATCH 1/3] feat: add a classification provider interface --- .env.example | 9 + librechat.example.yaml | 12 + packages/api/src/classification/index.ts | 5 + .../src/classification/providers/http.spec.ts | 295 ++++++++++++++++++ .../api/src/classification/providers/http.ts | 115 +++++++ .../src/classification/providers/transport.ts | 184 +++++++++++ .../api/src/classification/questions.spec.ts | 124 ++++++++ packages/api/src/classification/questions.ts | 82 +++++ packages/api/src/classification/registry.ts | 48 +++ packages/api/src/classification/resolve.ts | 72 +++++ packages/api/src/classification/types.ts | 126 ++++++++ packages/api/src/index.ts | 2 + packages/data-provider/src/config.ts | 24 ++ packages/data-schemas/src/app/service.ts | 2 + packages/data-schemas/src/types/app.ts | 2 + 15 files changed, 1102 insertions(+) create mode 100644 packages/api/src/classification/index.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..d53f458a0dd 100644 --- a/.env.example +++ b/.env.example @@ -1403,6 +1403,15 @@ 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 + #======================# # MCP Configuration # #======================# diff --git a/librechat.example.yaml b/librechat.example.yaml index 260912aa780..c921ff4d4b3 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -515,6 +515,18 @@ 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 + # Example MCP Servers Object Structure # mcpServers: # everything: 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/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..d5927d6411c --- /dev/null +++ b/packages/api/src/classification/providers/http.ts @@ -0,0 +1,115 @@ +import type { + Classifier, + ClassificationUsage, + ClassificationAnswer, + ClassificationResult, + ClassificationRequest, +} from '../types'; +import type { ProviderFetch } from './transport'; +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; + timeoutMs?: number; + maxRetries?: number; + fetch?: ProviderFetch; + sleep?: (ms: number) => Promise; +} + +function readAnswer(answer: unknown): ClassificationAnswer | null { + if (answer == null || typeof answer !== 'object') { + return null; + } + const record = answer as { + type?: unknown; + probability?: unknown; + choice?: unknown; + score?: unknown; + confidence?: unknown; + probabilities?: unknown; + }; + const probabilities = (record.probabilities ?? {}) as Record; + const confidence = typeof record.confidence === 'number' ? record.confidence : null; + + if (record.type === 'boolean' && typeof record.probability === 'number') { + return { type: 'boolean', probability: record.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; +} + +export function parseEnvelope( + body: string, + providerId: string, + readOne: (answer: unknown) => ClassificationAnswer | null, +): 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 record = parsed 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 ?? ''; + + return { + id: PROVIDER_ID, + model, + async classify(request: ClassificationRequest): Promise { + const payload = JSON.stringify({ + ...(model ? { model } : {}), + state: request.state, + questions: request.questions, + }); + const body = await send(payload, request.signal, request.label ?? 'classify'); + return parseEnvelope(body, PROVIDER_ID, readAnswer); + }, + }; +} diff --git a/packages/api/src/classification/providers/transport.ts b/packages/api/src/classification/providers/transport.ts new file mode 100644 index 00000000000..959103e9f56 --- /dev/null +++ b/packages/api/src/classification/providers/transport.ts @@ -0,0 +1,184 @@ +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, +) => 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 timeoutMs = 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): 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) { + let lastError: ClassificationError | undefined; + for (let attemptNo = 0; attemptNo <= maxRetries; attemptNo++) { + try { + const started = Date.now(); + const body = await attempt(payload, signal); + 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..b9d9b7223ff --- /dev/null +++ b/packages/api/src/classification/registry.ts @@ -0,0 +1,48 @@ +import type { ProviderFetch } from './providers/transport'; +import type { Classifier } from './types'; +import { createHttpClassifier, PROVIDER_ID as HTTP_ID } from './providers/http'; + +export interface ProviderSettings { + baseURL?: string; + model?: string; + timeoutMs?: number; + maxRetries?: number; + apiKeyEnv?: string; +} + +export interface ProviderBuildParams { + settings: ProviderSettings; + apiKey: string; + fetch?: ProviderFetch; +} + +export interface ProviderEntry { + create(params: ProviderBuildParams): Classifier; +} + +export const PROVIDERS: Record = { + [HTTP_ID]: { + create: ({ settings, apiKey, fetch }) => + createHttpClassifier({ + apiKey, + endpoint: settings.baseURL ?? '', + model: settings.model, + timeoutMs: settings.timeoutMs, + maxRetries: settings.maxRetries, + fetch, + }), + }, +}; + +export const DEFAULT_API_KEY_ENV = 'CLASSIFIER_API_KEY'; + +export function getProvider(id: string | undefined): ProviderEntry | null { + if (!id) { + return null; + } + return PROVIDERS[id] ?? null; +} + +export function providerNames(): string[] { + return Object.keys(PROVIDERS); +} diff --git a/packages/api/src/classification/resolve.ts b/packages/api/src/classification/resolve.ts new file mode 100644 index 00000000000..862cd71a08b --- /dev/null +++ b/packages/api/src/classification/resolve.ts @@ -0,0 +1,72 @@ +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 { getProvider, 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 entry = getProvider(providerId); + if (entry == null) { + if (!warned.has(`unknown:${providerId}`)) { + warned.add(`unknown:${providerId}`); + logger.warn( + `[classification] unknown provider "${providerId}"; this build supports: ` + + `${providerNames().join(', ')}. Classification stays off.`, + ); + } + return null; + } + + const settings = config.providers?.[providerId] ?? {}; + 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 = entry.create({ settings, apiKey, fetch: 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..7b3e352ceef --- /dev/null +++ b/packages/api/src/classification/types.ts @@ -0,0 +1,126 @@ +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; +} + +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..ff9287974a9 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2908,12 +2908,36 @@ 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(), + /** 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(), +}); + +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 e19fad36dc3b40de9feec3d936f44b00fd1895d3 Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Tue, 22 Sep 2026 13:04:28 +0900 Subject: [PATCH 2/3] 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 c921ff4d4b3..012a552537c 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -526,6 +526,30 @@ actions: # baseURL: https://classifier.example.com/v1/classify # apiKeyEnv: CLASSIFIER_API_KEY # timeoutMs: 4000 +# +# # 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 862cd71a08b..998424da8fe 100644 --- a/packages/api/src/classification/resolve.ts +++ b/packages/api/src/classification/resolve.ts @@ -70,3 +70,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 ff9287974a9..252e2437c89 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2927,6 +2927,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 ec01c3e2528395cebe78345c0fb86d0f8eb0617d Mon Sep 17 00:00:00 2001 From: Ali Gulzar Date: Tue, 22 Sep 2026 13:34:57 +0900 Subject: [PATCH 3/3] 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 012a552537c..f6e2dc58e72 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -550,6 +550,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 252e2437c89..cc7b888e81c 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2959,7 +2959,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 @@ -2967,6 +2966,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({}), });