diff --git a/.env.example b/.env.example index a85ba4c9234..569027064cb 100644 --- a/.env.example +++ b/.env.example @@ -1403,6 +1403,20 @@ OPENWEATHER_API_KEY= # or # COHERE_API_KEY=your_cohere_api_key +#======================# +# Classification # +#======================# + +# Key for the provider named by `classification.provider` in librechat.yaml. +# Each provider declares which variable it reads through +# `classification.providers..apiKeyEnv`; this is the default. +# CLASSIFIER_API_KEY=your_classifier_api_key + +# The presets read these instead. +# TYPESAFE_API_KEY=your_typesafe_api_key +# OPENROUTER_KEY=your_openrouter_key +# CLOUDFLARE_API_TOKEN=your_cloudflare_api_token + #======================# # MCP Configuration # #======================# diff --git a/librechat.example.yaml b/librechat.example.yaml index 260912aa780..5d159a89ef7 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -515,6 +515,39 @@ actions: # - 'host.docker.internal:8080' # - '127.0.0.1:8080' +# Classification: small typed judgments (yes/no, pick one, rate) that code can +# branch on. Off unless enabled. The key is read from the environment variable +# named by apiKeyEnv, never written here. +# classification: +# enabled: true +# provider: http +# providers: +# http: +# baseURL: https://classifier.example.com/v1/classify +# apiKeyEnv: CLASSIFIER_API_KEY +# timeoutMs: 4000 +# +# # Known hosts ship as presets, so naming one is usually enough. Cloudflare +# # is the exception: its URL carries your account id. +# provider: cloudflare +# providers: +# cloudflare: +# baseURL: https://api.cloudflare.com/client/v4/accounts//ai/run +# apiKeyEnv: CLOUDFLARE_API_TOKEN +# +# # A host with no preset needs no code either, only its shape. `dialect` +# # picks the wire vocabulary, `requestKey` nests the body, `responseKey` +# # unwraps the reply. +# provider: inhouse +# providers: +# inhouse: +# baseURL: https://classify.internal/v1/run +# model: your-model +# dialect: systemone +# requestKey: input +# responseKey: result +# apiKeyEnv: INHOUSE_CLASSIFIER_KEY + # Example MCP Servers Object Structure # mcpServers: # everything: diff --git a/packages/api/src/classification/config.spec.ts b/packages/api/src/classification/config.spec.ts new file mode 100644 index 00000000000..b1421305e79 --- /dev/null +++ b/packages/api/src/classification/config.spec.ts @@ -0,0 +1,134 @@ +import { classificationSchema } from 'librechat-data-provider'; +import type { TClassificationConfig } from 'librechat-data-provider'; +import { resolveClassifier } from './resolve'; +import { boolean } from './questions'; + +/** + * Covers the path an operator's `librechat.yaml` actually travels: the zod + * schema, then the resolver, then the bytes on the wire. A field that parses + * but never reaches the request is the failure this file exists to catch. + */ + +const ANSWER = { model: 'jev-1.13.0', answers: { d: { type: 'noul', noul: 0.8 } }, usage: {} }; + +function recorder(response: unknown = ANSWER) { + const calls: { url: string; body: Record }[] = []; + const fetch = async (url: string, init: { body?: string }) => { + calls.push({ url, body: JSON.parse(init.body ?? '{}') }); + return { + ok: true, + status: 200, + headers: { get: () => null }, + text: async () => JSON.stringify(response), + }; + }; + return { calls, fetch: fetch as never }; +} + +function parse(raw: unknown): TClassificationConfig { + return classificationSchema.parse(raw) as TClassificationConfig; +} + +describe('classification config', () => { + it('parses an unset block into every capability off', () => { + const config = parse({}); + + expect(config.enabled).toBe(false); + expect(config.provider).toBe('http'); + expect(config.providers).toEqual({}); + }); + + it('keeps the wire-shape fields through parsing', () => { + const config = parse({ + enabled: true, + provider: 'cloudflare', + providers: { + cloudflare: { + baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run', + model: 'typesafe/jev', + dialect: 'systemone', + requestKey: 'input', + responseKey: 'result', + apiKeyEnv: 'CLOUDFLARE_API_TOKEN', + timeoutMs: 9000, + maxRetries: 1, + }, + }, + }); + + expect(config.providers.cloudflare).toEqual({ + baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run', + model: 'typesafe/jev', + dialect: 'systemone', + requestKey: 'input', + responseKey: 'result', + apiKeyEnv: 'CLOUDFLARE_API_TOKEN', + timeoutMs: 9000, + maxRetries: 1, + }); + }); + + it('rejects a dialect it does not implement', () => { + const result = classificationSchema.safeParse({ + enabled: true, + provider: 'x', + providers: { x: { baseURL: 'https://x.test/run', dialect: 'logprobs' } }, + }); + + expect(result.success).toBe(false); + }); + + it('rejects a misspelled key instead of dropping it', () => { + const result = classificationSchema.safeParse({ + enabled: true, + provider: 'x', + providers: { x: { baseURL: 'https://x.test/run', requestkey: 'input' } }, + }); + + expect(result.success).toBe(false); + }); + + it('rejects a baseURL that is not a URL', () => { + const result = classificationSchema.safeParse({ + enabled: true, + provider: 'x', + providers: { x: { baseURL: 'classify.internal' } }, + }); + + expect(result.success).toBe(false); + }); + + it('carries a parsed config all the way onto the wire', async () => { + const config = parse({ + enabled: true, + provider: 'cloudflare', + providers: { + cloudflare: { baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run' }, + }, + }); + const { calls, fetch } = recorder({ result: ANSWER }); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://api.cloudflare.com/client/v4/accounts/abc/ai/run'); + expect(calls[0].body).toHaveProperty('input.questions.d.type', 'noul'); + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('lets a parsed override beat the preset it sits on', async () => { + const config = parse({ + enabled: true, + provider: 'typesafe', + providers: { typesafe: { baseURL: 'https://proxy.internal/systemone', model: 'jev-1.13' } }, + }); + const { calls, fetch } = recorder(); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://proxy.internal/systemone'); + expect(calls[0].body.model).toBe('jev-1.13'); + expect((calls[0].body.questions as Record).d.type).toBe('noul'); + }); +}); diff --git a/packages/api/src/classification/index.ts b/packages/api/src/classification/index.ts new file mode 100644 index 00000000000..fce04f418ea --- /dev/null +++ b/packages/api/src/classification/index.ts @@ -0,0 +1,5 @@ +export * from './types'; +export * from './questions'; +export * from './registry'; +export * from './resolve'; +export type { ProviderFetch, Transport, TransportOptions } from './providers/transport'; diff --git a/packages/api/src/classification/presets.spec.ts b/packages/api/src/classification/presets.spec.ts new file mode 100644 index 00000000000..0bfcd642ad4 --- /dev/null +++ b/packages/api/src/classification/presets.spec.ts @@ -0,0 +1,120 @@ +import type { TClassificationConfig } from 'librechat-data-provider'; +import { PRESETS, presetFor, mergeSettings } from './registry'; +import { resolveClassifier } from './resolve'; +import { boolean } from './questions'; + +type Captured = { url: string; body: Record }; + +function recorder(response: unknown) { + const calls: Captured[] = []; + const fetch = async (url: string, init: { body?: string }) => { + calls.push({ url, body: JSON.parse(init.body ?? '{}') }); + return { + ok: true, + status: 200, + headers: { get: () => null }, + text: async () => JSON.stringify(response), + }; + }; + return { calls, fetch: fetch as never }; +} + +function configFor(provider: string, settings?: Record): TClassificationConfig { + return { + enabled: true, + provider, + providers: settings == null ? {} : { [provider]: settings }, + toolSelection: {}, + memoryGate: {}, + } as unknown as TClassificationConfig; +} + +const NOUL = { model: 'jev-1.13.0', answers: { d: { type: 'noul', noul: 0.8 } }, usage: {} }; + +describe('presets', () => { + it('ships every known host as settings, not as code', () => { + expect(Object.keys(PRESETS).sort()).toEqual(['cloudflare', 'http', 'openrouter', 'typesafe']); + }); + + it('lets an operator override any field of a preset', () => { + const merged = mergeSettings(presetFor('typesafe'), { model: 'jev-1.13', timeoutMs: 9000 }); + + expect(merged.model).toBe('jev-1.13'); + expect(merged.timeoutMs).toBe(9000); + expect(merged.baseURL).toBe('https://api.typesafe.ai/v1/systemone'); + }); + + it('sends the System One vocabulary for the typesafe preset', async () => { + const { calls, fetch } = recorder(NOUL); + const classifier = resolveClassifier({ config: configFor('typesafe'), apiKey: 'k', fetch }); + + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://api.typesafe.ai/v1/systemone'); + expect((calls[0].body.questions as Record).d.type).toBe('noul'); + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('nests the body and unwraps the envelope for the cloudflare preset', async () => { + const { calls, fetch } = recorder({ result: NOUL, success: true }); + const config = configFor('cloudflare', { + baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run', + }); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].body).toHaveProperty('input.questions.d.type', 'noul'); + expect(calls[0].body.model).toBe('typesafe/jev'); + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('still reads a bare body when the host does not wrap it', async () => { + const { fetch } = recorder(NOUL); + const config = configFor('cloudflare', { + baseURL: 'https://api.cloudflare.com/client/v4/accounts/abc/ai/run', + }); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('sends a flat body for the openrouter preset', async () => { + const { calls, fetch } = recorder(NOUL); + const classifier = resolveClassifier({ config: configFor('openrouter'), apiKey: 'k', fetch }); + + await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://openrouter.ai/api/alpha/decisions'); + expect(calls[0].body).not.toHaveProperty('input'); + expect(calls[0].body.model).toBe('~typesafe/jev-latest'); + }); + + it('builds a host it has never heard of from config alone', async () => { + const { calls, fetch } = recorder({ data: NOUL }); + const config = configFor('somethingnew', { + baseURL: 'https://classify.internal/v1/run', + model: 'house-classifier', + dialect: 'systemone', + requestKey: 'payload', + responseKey: 'data', + }); + + const classifier = resolveClassifier({ config, apiKey: 'k', fetch }); + const result = await classifier!.classify({ state: {}, questions: { d: boolean('durable?') } }); + + expect(calls[0].url).toBe('https://classify.internal/v1/run'); + expect(calls[0].body).toHaveProperty('payload.questions.d.type', 'noul'); + expect(result.answers.d).toEqual({ type: 'boolean', probability: 0.8 }); + }); + + it('stays off for an unknown name with no baseURL to go on', () => { + expect(resolveClassifier({ config: configFor('mystery'), apiKey: 'k' })).toBeNull(); + }); + + it('stays off when a preset needs a baseURL the operator did not give', () => { + expect(resolveClassifier({ config: configFor('cloudflare'), apiKey: 'k' })).toBeNull(); + }); +}); diff --git a/packages/api/src/classification/providers/dialect.spec.ts b/packages/api/src/classification/providers/dialect.spec.ts new file mode 100644 index 00000000000..eaeb4c937ab --- /dev/null +++ b/packages/api/src/classification/providers/dialect.spec.ts @@ -0,0 +1,65 @@ +import { toWireQuestion, readAnswer } from './dialect'; +import { boolean, choice, score } from '../questions'; + +describe('toWireQuestion', () => { + it('keeps the port vocabulary by default', () => { + expect(toWireQuestion(boolean('durable?'), 'port').type).toBe('boolean'); + }); + + it('renames a yes/no question for System One', () => { + expect(toWireQuestion(boolean('durable?'), 'systemone').type).toBe('noul'); + }); + + it('leaves choice and score alone in either dialect', () => { + const pick = choice('which', { a: null, b: null }); + const rate = score('how much', ['low', 'high']); + + expect(toWireQuestion(pick, 'systemone').type).toBe('choice'); + expect(toWireQuestion(rate, 'systemone').type).toBe('score'); + expect(toWireQuestion(pick, 'port').criteria).toEqual({ a: null, b: null }); + }); + + it('omits criteria when the question carries none', () => { + expect(toWireQuestion(boolean('durable?'), 'port')).not.toHaveProperty('criteria'); + }); +}); + +describe('readAnswer', () => { + it('reads a port boolean', () => { + expect(readAnswer({ type: 'boolean', probability: 0.9 }, 'port')).toEqual({ + type: 'boolean', + probability: 0.9, + }); + }); + + it('reads a System One noul as a boolean', () => { + expect(readAnswer({ type: 'noul', noul: 0.9 }, 'systemone')).toEqual({ + type: 'boolean', + probability: 0.9, + }); + }); + + it('does not read a noul when the port dialect was asked for', () => { + expect(readAnswer({ type: 'noul', noul: 0.9 }, 'port')).toBeNull(); + }); + + it('reports an unmeasured confidence as null rather than zero', () => { + const answer = readAnswer( + { type: 'choice', choice: 'a', probabilities: { a: 1 } }, + 'systemone', + ); + + expect(answer).toEqual({ + type: 'choice', + choice: 'a', + confidence: null, + probabilities: { a: 1 }, + }); + }); + + it('drops an answer it cannot read rather than inventing one', () => { + expect(readAnswer({ type: 'choice' }, 'port')).toBeNull(); + expect(readAnswer(null, 'port')).toBeNull(); + expect(readAnswer('yes', 'port')).toBeNull(); + }); +}); diff --git a/packages/api/src/classification/providers/dialect.ts b/packages/api/src/classification/providers/dialect.ts new file mode 100644 index 00000000000..1f267a290b0 --- /dev/null +++ b/packages/api/src/classification/providers/dialect.ts @@ -0,0 +1,49 @@ +import type { ClassificationAnswer, ClassificationQuestion } from '../types'; + +export type Dialect = 'port' | 'systemone'; + +interface WireQuestion { + type: string; + instructions: unknown; + criteria?: unknown; +} + +interface WireAnswer { + type?: unknown; + probability?: unknown; + noul?: unknown; + choice?: unknown; + score?: unknown; + confidence?: unknown; + probabilities?: unknown; +} + +/** System One calls a yes/no question a `noul`; the port calls it a boolean. */ +export function toWireQuestion(question: ClassificationQuestion, dialect: Dialect): WireQuestion { + const type = dialect === 'systemone' && question.type === 'boolean' ? 'noul' : question.type; + return question.criteria == null + ? { type, instructions: question.instructions } + : { type, instructions: question.instructions, criteria: question.criteria }; +} + +export function readAnswer(answer: unknown, dialect: Dialect): ClassificationAnswer | null { + if (answer == null || typeof answer !== 'object') { + return null; + } + const record = answer as WireAnswer; + const probabilities = (record.probabilities ?? {}) as Record; + const confidence = typeof record.confidence === 'number' ? record.confidence : null; + + const booleanType = dialect === 'systemone' ? 'noul' : 'boolean'; + const probability = dialect === 'systemone' ? record.noul : record.probability; + if (record.type === booleanType && typeof probability === 'number') { + return { type: 'boolean', probability }; + } + if (record.type === 'choice' && typeof record.choice === 'string') { + return { type: 'choice', choice: record.choice, confidence, probabilities }; + } + if (record.type === 'score' && typeof record.score === 'number') { + return { type: 'score', score: record.score, confidence, probabilities }; + } + return null; +} diff --git a/packages/api/src/classification/providers/http.spec.ts b/packages/api/src/classification/providers/http.spec.ts new file mode 100644 index 00000000000..6234637ed4e --- /dev/null +++ b/packages/api/src/classification/providers/http.spec.ts @@ -0,0 +1,295 @@ +import type { ProviderFetch } from './transport'; +import { ClassificationError } from '../types'; +import { createHttpClassifier } from './http'; + +interface StubResponse { + ok: boolean; + status: number; + body: string; + headers?: Record; +} + +function stubTransport(queue: Array): { + transport: ProviderFetch; + calls: Array<{ url: string; headers: Record; body: string }>; +} { + const calls: Array<{ url: string; headers: Record; body: string }> = []; + let index = 0; + const transport: ProviderFetch = async (url, init) => { + calls.push({ url, headers: init.headers, body: init.body }); + const next = queue[Math.min(index, queue.length - 1)]; + index++; + if (next instanceof Error) { + throw next; + } + return { + ok: next.ok, + status: next.status, + headers: { get: (name: string) => next.headers?.[name.toLowerCase()] ?? null }, + text: async () => next.body, + }; + }; + return { transport, calls }; +} + +const ENDPOINT = 'https://classifier.test/v1/classify'; + +const ANSWER = JSON.stringify({ + model: 'test-1', + answers: { verdict: { type: 'boolean', probability: 0.82 } }, + usage: { input_tokens: 120, output_tokens: 8 }, +}); + +const QUESTION = { + verdict: { type: 'boolean' as const, instructions: 'Is this urgent?' }, +}; + +function build(queue: Array, overrides = {}) { + const { transport, calls } = stubTransport(queue); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + fetch: transport, + sleep: async () => undefined, + ...overrides, + }); + return { classifier, calls }; +} + +describe('createHttpClassifier', () => { + it('posts the state and questions and reads the answers back', async () => { + const { classifier, calls } = build([{ ok: true, status: 200, body: ANSWER }]); + + const result = await classifier.classify({ state: 'payouts are failing', questions: QUESTION }); + + expect(result.answers.verdict).toEqual({ type: 'boolean', probability: 0.82 }); + expect(result.usage).toEqual({ inputTokens: 120, outputTokens: 8 }); + expect(calls[0].url).toBe(ENDPOINT); + expect(calls[0].headers.Authorization).toBe('Bearer sk-test'); + expect(JSON.parse(calls[0].body)).toEqual({ + state: 'payouts are failing', + questions: QUESTION, + }); + }); + + it('includes the model only when one is configured', async () => { + const { classifier, calls } = build([{ ok: true, status: 200, body: ANSWER }], { + model: 'test-1', + }); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(JSON.parse(calls[0].body).model).toBe('test-1'); + expect(classifier.model).toBe('test-1'); + expect(classifier.id).toBe('http'); + }); + + it('carries choice and score answers through', async () => { + const body = JSON.stringify({ + model: 'test-1', + answers: { + pick: { type: 'choice', choice: 'b', confidence: 0.7, probabilities: { a: 0.3, b: 0.7 } }, + rate: { type: 'score', score: 1.4, confidence: 0.5, probabilities: { '0': 0.6, '1': 0.4 } }, + }, + usage: { input_tokens: 10, output_tokens: 2 }, + }); + const { classifier } = build([{ ok: true, status: 200, body }]); + + const result = await classifier.classify({ + state: 'x', + questions: { + pick: { type: 'choice', instructions: 'which', criteria: { a: null, b: null } }, + rate: { type: 'score', instructions: 'how much', criteria: ['low', 'high'] }, + }, + }); + + expect(result.answers.pick).toMatchObject({ type: 'choice', choice: 'b', confidence: 0.7 }); + expect(result.answers.rate).toMatchObject({ type: 'score', score: 1.4 }); + }); + + it('retries a 429 and honors retry-after', async () => { + const waits: number[] = []; + const { transport, calls } = stubTransport([ + { ok: false, status: 429, body: 'slow down', headers: { 'retry-after': '2' } }, + { ok: true, status: 200, body: ANSWER }, + ]); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + fetch: transport, + sleep: async (ms) => { + waits.push(ms); + }, + }); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(calls).toHaveLength(2); + expect(waits).toEqual([2000]); + }); + + it('clamps an absurd retry-after', async () => { + const waits: number[] = []; + const { transport } = stubTransport([ + { ok: false, status: 503, body: 'down', headers: { 'retry-after': '3600' } }, + { ok: true, status: 200, body: ANSWER }, + ]); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + fetch: transport, + sleep: async (ms) => { + waits.push(ms); + }, + }); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(waits).toEqual([10_000]); + }); + + it('gives up after maxRetries and names the provider', async () => { + const { classifier, calls } = build([{ ok: false, status: 500, body: 'boom' }], { + maxRetries: 2, + }); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + name: 'ClassificationError', + failure: 'server_error', + provider: 'http', + status: 500, + }); + expect(calls).toHaveLength(3); + }); + + it('does not retry a rejected request', async () => { + const { classifier, calls } = build([{ ok: false, status: 422, body: 'bad question' }]); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'bad_request', + }); + expect(calls).toHaveLength(1); + }); + + it('does not retry a rejected key', async () => { + const { classifier, calls } = build([{ ok: false, status: 401, body: 'nope' }]); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'unauthorized', + }); + expect(calls).toHaveLength(1); + }); + + it('reports a body that is not JSON as malformed', async () => { + const { classifier } = build([{ ok: true, status: 200, body: '504' }]); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'malformed_response', + }); + }); + + it('reports a JSON body with no answers as malformed', async () => { + const { classifier } = build([ + { ok: true, status: 200, body: JSON.stringify({ model: 'test-1' }) }, + ]); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'malformed_response', + }); + }); + + it('drops an answer it cannot read rather than inventing one', async () => { + const body = JSON.stringify({ + model: 'test-1', + answers: { verdict: { type: 'something_new', value: 3 } }, + usage: { input_tokens: 1, output_tokens: 1 }, + }); + const { classifier } = build([{ ok: true, status: 200, body }]); + + const result = await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(result.answers.verdict).toBeUndefined(); + }); + + it('defaults usage when the response omits it', async () => { + const { classifier } = build([ + { + ok: true, + status: 200, + body: JSON.stringify({ + model: 'test-1', + answers: { verdict: { type: 'boolean', probability: 1 } }, + }), + }, + ]); + + const result = await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(result.usage).toEqual({ inputTokens: 0, outputTokens: 0 }); + }); + + it('times out a transport that never settles', async () => { + const transport: ProviderFetch = (_url, init) => + new Promise((_resolve, reject) => { + init.signal.addEventListener('abort', () => reject(new Error('aborted'))); + }); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + timeoutMs: 15, + maxRetries: 0, + fetch: transport, + }); + + await expect(classifier.classify({ state: 'x', questions: QUESTION })).rejects.toMatchObject({ + failure: 'timeout', + }); + }); + + it('reports a caller abort as aborted, not as a timeout', async () => { + const controller = new AbortController(); + const transport: ProviderFetch = (_url, init) => + new Promise((_resolve, reject) => { + init.signal.addEventListener('abort', () => reject(new Error('aborted'))); + }); + const classifier = createHttpClassifier({ + apiKey: 'sk-test', + endpoint: ENDPOINT, + timeoutMs: 5_000, + maxRetries: 0, + fetch: transport, + }); + + const pending = classifier.classify({ + state: 'x', + questions: QUESTION, + signal: controller.signal, + }); + controller.abort(); + + await expect(pending).rejects.toMatchObject({ failure: 'aborted' }); + }); + + it('retries a transport-level network failure', async () => { + const { classifier, calls } = build([ + new Error('ECONNRESET'), + { ok: true, status: 200, body: ANSWER }, + ]); + + await classifier.classify({ state: 'x', questions: QUESTION }); + + expect(calls).toHaveLength(2); + }); + + it('refuses to build without a key', () => { + expect(() => createHttpClassifier({ apiKey: ' ', endpoint: ENDPOINT })).toThrow( + ClassificationError, + ); + }); + + it('refuses to build without an endpoint', () => { + expect(() => createHttpClassifier({ apiKey: 'sk-test', endpoint: '' })).toThrow( + ClassificationError, + ); + }); +}); diff --git a/packages/api/src/classification/providers/http.ts b/packages/api/src/classification/providers/http.ts new file mode 100644 index 00000000000..407ad65adad --- /dev/null +++ b/packages/api/src/classification/providers/http.ts @@ -0,0 +1,112 @@ +import type { + Classifier, + ClassificationUsage, + ClassificationAnswer, + ClassificationResult, + ClassificationRequest, +} from '../types'; +import type { ProviderFetch } from './transport'; +import type { Dialect } from './dialect'; +import { toWireQuestion, readAnswer } from './dialect'; +import { ClassificationError } from '../types'; +import { createTransport } from './transport'; + +export const PROVIDER_ID = 'http'; + +export interface HttpProviderOptions { + apiKey: string; + /** Full URL, not a base path. */ + endpoint: string; + model?: string; + /** Which wire vocabulary the endpoint speaks. */ + dialect?: Dialect; + /** Nests `state` and `questions` under this key, for hosts that wrap them. */ + requestKey?: string; + /** Reads the answer envelope from this key, for hosts that wrap the response. */ + responseKey?: string; + timeoutMs?: number; + maxRetries?: number; + fetch?: ProviderFetch; + sleep?: (ms: number) => Promise; +} + +export function parseEnvelope( + body: string, + providerId: string, + readOne: (answer: unknown) => ClassificationAnswer | null, + responseKey?: string, +): ClassificationResult { + let parsed: unknown; + try { + parsed = JSON.parse(body); + } catch { + throw new ClassificationError('malformed_response', 'response was not JSON', { + provider: providerId, + }); + } + if (parsed == null || typeof parsed !== 'object') { + throw new ClassificationError('malformed_response', 'response was not an object', { + provider: providerId, + }); + } + const unwrapped = + responseKey != null && responseKey !== '' + ? ((parsed as Record)[responseKey] ?? parsed) + : parsed; + const record = unwrapped as { model?: unknown; answers?: unknown; usage?: unknown }; + if (record.answers == null || typeof record.answers !== 'object') { + throw new ClassificationError('malformed_response', 'response carried no answers', { + provider: providerId, + }); + } + + const answers: Record = {}; + for (const [id, answer] of Object.entries(record.answers as Record)) { + const mapped = readOne(answer); + if (mapped != null) { + answers[id] = mapped; + } + } + + const raw = (record.usage ?? {}) as { input_tokens?: number; output_tokens?: number }; + const usage: ClassificationUsage = { + inputTokens: raw.input_tokens ?? 0, + outputTokens: raw.output_tokens ?? 0, + }; + + return { + model: typeof record.model === 'string' ? record.model : 'unknown', + answers, + usage, + }; +} + +export function createHttpClassifier(options: HttpProviderOptions): Classifier { + const send = createTransport({ providerId: PROVIDER_ID, ...options }); + const model = options.model ?? ''; + const dialect: Dialect = options.dialect ?? 'port'; + const { requestKey, responseKey } = options; + + return { + id: PROVIDER_ID, + model, + async classify(request: ClassificationRequest): Promise { + const questions: Record = {}; + for (const [id, question] of Object.entries(request.questions)) { + questions[id] = toWireQuestion(question, dialect); + } + const inner = { state: request.state, questions }; + const payload = JSON.stringify({ + ...(model ? { model } : {}), + ...(requestKey != null && requestKey !== '' ? { [requestKey]: inner } : inner), + }); + const body = await send( + payload, + request.signal, + request.label ?? 'classify', + request.timeoutMs, + ); + return parseEnvelope(body, PROVIDER_ID, (a) => readAnswer(a, dialect), responseKey); + }, + }; +} diff --git a/packages/api/src/classification/providers/transport.ts b/packages/api/src/classification/providers/transport.ts new file mode 100644 index 00000000000..c18a5302cca --- /dev/null +++ b/packages/api/src/classification/providers/transport.ts @@ -0,0 +1,191 @@ +import { logger } from '@librechat/data-schemas'; +import { ClassificationError } from '../types'; + +const DEFAULT_TIMEOUT_MS = 4_000; +const DEFAULT_MAX_RETRIES = 2; +const BACKOFF_MS = [250, 750, 1_500, 3_000, 6_000]; +const MAX_RETRY_AFTER_MS = 10_000; + +export type ProviderFetch = ( + input: string, + init: { + method: string; + headers: Record; + body: string; + signal: AbortSignal; + }, +) => Promise<{ + ok: boolean; + status: number; + headers: { get(name: string): string | null }; + text(): Promise; +}>; + +export interface TransportOptions { + providerId: string; + apiKey: string; + /** Full URL, not a base path. */ + endpoint: string; + timeoutMs?: number; + maxRetries?: number; + fetch?: ProviderFetch; + sleep?: (ms: number) => Promise; +} + +export type Transport = ( + payload: string, + signal: AbortSignal | undefined, + label: string, + timeoutOverrideMs?: number, +) => Promise; + +function failureForStatus(status: number) { + if (status === 401 || status === 403) { + return 'unauthorized' as const; + } + if (status === 429) { + return 'rate_limited' as const; + } + if (status >= 500) { + return 'server_error' as const; + } + return 'bad_request' as const; +} + +function isRetryable(failure: string): boolean { + return failure === 'rate_limited' || failure === 'server_error' || failure === 'network'; +} + +function retryAfterMs(header: string | null): number | undefined { + if (!header) { + return undefined; + } + const seconds = Number(header); + if (Number.isFinite(seconds) && seconds >= 0) { + return Math.min(seconds * 1_000, MAX_RETRY_AFTER_MS); + } + const at = Date.parse(header); + if (Number.isNaN(at)) { + return undefined; + } + return Math.min(Math.max(at - Date.now(), 0), MAX_RETRY_AFTER_MS); +} + +function briefly(body: string): string { + const flat = body.replace(/\s+/g, ' ').trim(); + return flat.length > 200 ? `${flat.slice(0, 200)}…` : flat; +} + +const defaultSleep = (ms: number): Promise => + new Promise((resolve) => setTimeout(resolve, ms)); + +export function createTransport(options: TransportOptions): Transport { + const { providerId } = options; + const apiKey = options.apiKey?.trim(); + if (!apiKey) { + throw new ClassificationError('unauthorized', 'classifier requires an API key', { + provider: providerId, + }); + } + + const endpoint = options.endpoint?.trim(); + if (!endpoint) { + throw new ClassificationError('bad_request', 'classifier requires a baseURL', { + provider: providerId, + }); + } + + const defaultTimeoutMs = options.timeoutMs ?? DEFAULT_TIMEOUT_MS; + const maxRetries = options.maxRetries ?? DEFAULT_MAX_RETRIES; + const sleep = options.sleep ?? defaultSleep; + const candidate = options.fetch ?? (globalThis.fetch as unknown as ProviderFetch | undefined); + if (typeof candidate !== 'function') { + throw new ClassificationError('network', 'no fetch implementation available', { + provider: providerId, + }); + } + const fetchImpl: ProviderFetch = candidate; + + async function attempt( + payload: string, + signal: AbortSignal | undefined, + timeoutMs: number, + ): Promise { + const timeout = new AbortController(); + const timer = setTimeout(() => timeout.abort(), timeoutMs); + const combined = signal != null ? AbortSignal.any([signal, timeout.signal]) : timeout.signal; + + try { + const response = await fetchImpl(endpoint, { + method: 'POST', + headers: { + Authorization: `Bearer ${apiKey}`, + 'Content-Type': 'application/json', + }, + body: payload, + signal: combined, + }); + + const text = await response.text(); + if (!response.ok) { + const error = new ClassificationError( + failureForStatus(response.status), + `classifier returned ${response.status}: ${briefly(text)}`, + { provider: providerId, status: response.status }, + ); + const wait = retryAfterMs(response.headers.get('retry-after')); + if (wait != null) { + error.retryAfterMs = wait; + } + throw error; + } + return text; + } catch (error) { + if (error instanceof ClassificationError) { + throw error; + } + if (signal?.aborted === true) { + throw new ClassificationError('aborted', 'caller aborted the request', { + provider: providerId, + }); + } + if (timeout.signal.aborted) { + throw new ClassificationError('timeout', `no answer within ${timeoutMs}ms`, { + provider: providerId, + }); + } + const message = error instanceof Error ? error.message : String(error); + throw new ClassificationError('network', message, { provider: providerId }); + } finally { + clearTimeout(timer); + } + } + + return async function send(payload, signal, label, timeoutOverrideMs) { + const timeoutMs = + timeoutOverrideMs != null && timeoutOverrideMs > 0 ? timeoutOverrideMs : defaultTimeoutMs; + let lastError: ClassificationError | undefined; + for (let attemptNo = 0; attemptNo <= maxRetries; attemptNo++) { + try { + const started = Date.now(); + const body = await attempt(payload, signal, timeoutMs); + logger.debug(`[classification] ${label} answered in ${Date.now() - started}ms`); + return body; + } catch (error) { + lastError = + error instanceof ClassificationError + ? error + : new ClassificationError('network', String(error), { provider: providerId }); + if (attemptNo === maxRetries || !isRetryable(lastError.failure)) { + break; + } + await sleep( + lastError.retryAfterMs ?? BACKOFF_MS[Math.min(attemptNo, BACKOFF_MS.length - 1)], + ); + } + } + throw ( + lastError ?? new ClassificationError('network', 'request failed', { provider: providerId }) + ); + }; +} diff --git a/packages/api/src/classification/questions.spec.ts b/packages/api/src/classification/questions.spec.ts new file mode 100644 index 00000000000..7551b1346cc --- /dev/null +++ b/packages/api/src/classification/questions.spec.ts @@ -0,0 +1,124 @@ +import type { ScoreAnswer, ChoiceAnswer, BooleanAnswer } from './types'; +import { + score, + level, + label, + choice, + isTrue, + ranked, + boolean, + normalized, + probabilityOf, +} from './questions'; + +describe('question builders', () => { + it('builds a boolean without criteria', () => { + expect(boolean('Is this urgent?')).toEqual({ + type: 'boolean', + instructions: 'Is this urgent?', + }); + }); + + it('builds a boolean with criteria', () => { + expect(boolean('Is this urgent?', { true: 'yes means', false: 'no means' })).toEqual({ + type: 'boolean', + instructions: 'Is this urgent?', + criteria: { true: 'yes means', false: 'no means' }, + }); + }); + + it('builds a choice and a score', () => { + expect(choice('Which team?', { billing: 'money', tech: 'bugs' })).toEqual({ + type: 'choice', + instructions: 'Which team?', + criteria: { billing: 'money', tech: 'bugs' }, + }); + expect(score('How angry?', ['Calm', 'Cross', 'Furious'])).toEqual({ + type: 'score', + instructions: 'How angry?', + criteria: ['Calm', 'Cross', 'Furious'], + }); + }); +}); + +describe('isTrue', () => { + const answer: BooleanAnswer = { type: 'boolean', probability: 0.8 }; + + it('compares against the threshold inclusively', () => { + expect(isTrue(answer, 0.8)).toBe(true); + expect(isTrue(answer, 0.81)).toBe(false); + }); + + it('is false for a missing answer', () => { + expect(isTrue(undefined, 0)).toBe(false); + }); +}); + +describe('choice helpers', () => { + const answer: ChoiceAnswer = { + type: 'choice', + choice: 'tech', + confidence: 0.7, + probabilities: { tech: 0.7, billing: 0.2, sales: 0.05 }, + }; + + it('reads one option probability', () => { + expect(probabilityOf(answer, 'billing')).toBeCloseTo(0.2); + }); + + it('returns zero for an option that was never offered', () => { + expect(probabilityOf(answer, 'absent')).toBe(0); + }); + + it('ranks options most probable first', () => { + expect(ranked(answer)).toEqual(['tech', 'billing', 'sales']); + }); + + it('drops options under the floor', () => { + expect(ranked(answer, 0.1)).toEqual(['tech', 'billing']); + }); + + it('handles a missing answer', () => { + expect(ranked(undefined)).toEqual([]); + }); +}); + +describe('score helpers', () => { + const levels = ['Calm', 'Frustrated', 'Very angry']; + const answer: ScoreAnswer = { + type: 'score', + score: 1.24, + confidence: 0.6, + probabilities: { '0': 0.12, '1': 0.52, '2': 0.36 }, + }; + + it('reports the most probable level, not the rounded score', () => { + expect(level(answer)).toBe(1); + }); + + it('labels that level', () => { + expect(label(answer, levels)).toBe('Frustrated'); + }); + + it('returns an empty label when the levels do not cover it', () => { + expect(label(answer, ['only one'])).toBe(''); + }); + + it('normalizes the weighted score against the top level', () => { + expect(normalized(answer, levels.length)).toBeCloseTo(0.62); + }); + + it('clamps a score outside the level range', () => { + expect(normalized({ ...answer, score: 99 }, levels.length)).toBe(1); + expect(normalized({ ...answer, score: -1 }, levels.length)).toBe(0); + }); + + it('returns zero when there are too few levels to normalize', () => { + expect(normalized(answer, 1)).toBe(0); + }); + + it('handles a missing answer', () => { + expect(level(undefined)).toBe(0); + expect(normalized(undefined, 3)).toBe(0); + }); +}); diff --git a/packages/api/src/classification/questions.ts b/packages/api/src/classification/questions.ts new file mode 100644 index 00000000000..471ffe836b8 --- /dev/null +++ b/packages/api/src/classification/questions.ts @@ -0,0 +1,82 @@ +import type { + ScoreAnswer, + ChoiceAnswer, + ScoreQuestion, + BooleanAnswer, + ChoiceQuestion, + BooleanQuestion, + ClassificationText, +} from './types'; + +export function boolean( + instructions: ClassificationText, + criteria?: { true?: ClassificationText; false?: ClassificationText }, +): BooleanQuestion { + return criteria == null + ? { type: 'boolean', instructions } + : { type: 'boolean', instructions, criteria }; +} + +export function choice( + instructions: ClassificationText, + criteria: Record, +): ChoiceQuestion { + return { type: 'choice', instructions, criteria }; +} + +export function score( + instructions: ClassificationText, + levels: ClassificationText[], +): ScoreQuestion { + return { type: 'score', instructions, criteria: levels }; +} + +export function isTrue(answer: BooleanAnswer | undefined, threshold: number): boolean { + return answer != null && answer.probability >= threshold; +} + +export function probabilityOf( + answer: ChoiceAnswer | ScoreAnswer | undefined, + option: string, +): number { + return answer?.probabilities?.[option] ?? 0; +} + +/** Options above `floor`, most probable first. */ +export function ranked(answer: ChoiceAnswer | undefined, floor = 0): string[] { + if (answer == null) { + return []; + } + return Object.entries(answer.probabilities) + .filter(([, probability]) => probability >= floor) + .sort((a, b) => b[1] - a[1]) + .map(([option]) => option); +} + +/** The most probable level, which is not always the rounded weighted score. */ +export function level(answer: ScoreAnswer | undefined): number { + if (answer == null) { + return 0; + } + let best = 0; + let bestProbability = -1; + for (const [key, probability] of Object.entries(answer.probabilities)) { + if (probability > bestProbability) { + bestProbability = probability; + best = Number(key); + } + } + return Number.isFinite(best) ? best : 0; +} + +export function label(answer: ScoreAnswer | undefined, levels: readonly string[]): string { + return levels[level(answer)] ?? ''; +} + +/** The weighted score as a fraction of the highest level. */ +export function normalized(answer: ScoreAnswer | undefined, levelCount: number): number { + if (answer == null || levelCount < 2) { + return 0; + } + return Math.min(Math.max(answer.score / (levelCount - 1), 0), 1); +} diff --git a/packages/api/src/classification/registry.ts b/packages/api/src/classification/registry.ts new file mode 100644 index 00000000000..65e2ec6a5ae --- /dev/null +++ b/packages/api/src/classification/registry.ts @@ -0,0 +1,87 @@ +import type { ProviderFetch } from './providers/transport'; +import type { Dialect } from './providers/dialect'; +import type { Classifier } from './types'; +import { createHttpClassifier } from './providers/http'; + +export interface ProviderSettings { + baseURL?: string; + model?: string; + dialect?: Dialect; + requestKey?: string; + responseKey?: string; + timeoutMs?: number; + maxRetries?: number; + apiKeyEnv?: string; +} + +export const DEFAULT_API_KEY_ENV = 'CLASSIFIER_API_KEY'; + +/** + * Known hosts, as settings rather than code. They all serve the same question + * shapes over HTTP and differ only in URL, model name and how the body is + * wrapped, so a new one is an entry here or, for an operator who cannot wait + * for a release, the same fields written in `librechat.yaml`. + */ +export const PRESETS: Record = { + http: { + dialect: 'port', + apiKeyEnv: DEFAULT_API_KEY_ENV, + }, + typesafe: { + baseURL: 'https://api.typesafe.ai/v1/systemone', + model: 'jev-latest', + dialect: 'systemone', + apiKeyEnv: 'TYPESAFE_API_KEY', + }, + openrouter: { + baseURL: 'https://openrouter.ai/api/alpha/decisions', + model: '~typesafe/jev-latest', + dialect: 'systemone', + apiKeyEnv: 'OPENROUTER_KEY', + }, + cloudflare: { + /** No default URL: the account id is part of it. */ + model: 'typesafe/jev', + dialect: 'systemone', + requestKey: 'input', + responseKey: 'result', + apiKeyEnv: 'CLOUDFLARE_API_TOKEN', + }, +}; + +export function presetFor(name: string | undefined): ProviderSettings | null { + if (!name) { + return null; + } + return PRESETS[name] ?? null; +} + +/** Operator settings win over the preset, field by field. */ +export function mergeSettings( + preset: ProviderSettings | null, + configured: ProviderSettings | undefined, +): ProviderSettings { + return { ...(preset ?? {}), ...(configured ?? {}) }; +} + +export function providerNames(): string[] { + return Object.keys(PRESETS); +} + +export function createClassifier( + settings: ProviderSettings, + apiKey: string, + fetch?: ProviderFetch, +): Classifier { + return createHttpClassifier({ + apiKey, + endpoint: settings.baseURL ?? '', + model: settings.model, + dialect: settings.dialect, + requestKey: settings.requestKey, + responseKey: settings.responseKey, + timeoutMs: settings.timeoutMs, + maxRetries: settings.maxRetries, + fetch, + }); +} diff --git a/packages/api/src/classification/resolve.ts b/packages/api/src/classification/resolve.ts new file mode 100644 index 00000000000..a5734e4ad81 --- /dev/null +++ b/packages/api/src/classification/resolve.ts @@ -0,0 +1,89 @@ +import { logger } from '@librechat/data-schemas'; +import type { TClassificationConfig } from 'librechat-data-provider'; +import type { ProviderFetch } from './providers/transport'; +import type { Classifier } from './types'; +import { + presetFor, + mergeSettings, + createClassifier, + providerNames, + DEFAULT_API_KEY_ENV, +} from './registry'; +import { ClassificationError } from './types'; + +interface CacheEntry { + apiKey: string; + provider: string; + classifier: Classifier; +} + +const cache = new WeakMap(); +const warned = new Set(); + +export interface ResolveClassifierParams { + config?: TClassificationConfig | null; + apiKey?: string; + fetch?: ProviderFetch; +} + +export function resolveClassifier(params: ResolveClassifierParams): Classifier | null { + const config = params.config; + if (config == null || config.enabled !== true) { + return null; + } + + const providerId = config.provider; + const configured = config.providers?.[providerId]; + const preset = presetFor(providerId); + /** An unknown name is fine when the operator described the host in full. */ + if (preset == null && configured?.baseURL == null) { + if (!warned.has(`unknown:${providerId}`)) { + warned.add(`unknown:${providerId}`); + logger.warn( + `[classification] provider "${providerId}" has no preset and no baseURL. ` + + `Known presets: ${providerNames().join(', ')}. Classification stays off.`, + ); + } + return null; + } + + const settings = mergeSettings(preset, configured); + if (!settings.baseURL) { + if (!warned.has(`url:${providerId}`)) { + warned.add(`url:${providerId}`); + logger.warn( + `[classification] provider "${providerId}" needs a baseURL; classification stays off.`, + ); + } + return null; + } + const apiKeyEnv = settings.apiKeyEnv ?? DEFAULT_API_KEY_ENV; + const apiKey = (params.apiKey ?? process.env[apiKeyEnv] ?? '').trim(); + if (!apiKey) { + if (!warned.has(`key:${providerId}`)) { + warned.add(`key:${providerId}`); + logger.warn( + `[classification] provider "${providerId}" is configured but ${apiKeyEnv} is not set; ` + + 'every classification capability stays off.', + ); + } + return null; + } + + const cached = cache.get(config); + if (cached != null && cached.apiKey === apiKey && cached.provider === providerId) { + return cached.classifier; + } + + try { + const classifier = createClassifier(settings, apiKey, params.fetch); + cache.set(config, { apiKey, provider: providerId, classifier }); + return classifier; + } catch (error) { + logger.error( + `[classification] could not build provider "${providerId}"; capabilities stay off`, + error instanceof ClassificationError ? { failure: error.failure } : error, + ); + return null; + } +} diff --git a/packages/api/src/classification/types.ts b/packages/api/src/classification/types.ts new file mode 100644 index 00000000000..9a3cc9ae7da --- /dev/null +++ b/packages/api/src/classification/types.ts @@ -0,0 +1,128 @@ +export type ClassificationJson = + | string + | number + | boolean + | null + | ClassificationJson[] + | { [key: string]: ClassificationJson }; + +export type ClassificationText = + | string + | ClassificationJson[] + | { [key: string]: ClassificationJson }; + +export type ClassificationState = ClassificationText; + +export interface BooleanQuestion { + type: 'boolean'; + instructions: ClassificationText; + criteria?: { + true?: ClassificationText; + false?: ClassificationText; + }; +} + +export interface ChoiceQuestion { + type: 'choice'; + instructions: ClassificationText; + criteria: Record; +} + +export interface ScoreQuestion { + type: 'score'; + instructions: ClassificationText; + criteria: ClassificationText[]; +} + +export type ClassificationQuestion = BooleanQuestion | ChoiceQuestion | ScoreQuestion; + +export interface BooleanAnswer { + type: 'boolean'; + probability: number; +} + +export interface ChoiceAnswer { + type: 'choice'; + choice: string; + /** `null` when the provider cannot measure it, which is not the same as 0. */ + confidence: number | null; + probabilities: Record; +} + +export interface ScoreAnswer { + type: 'score'; + /** 0 to levels - 1. */ + score: number; + confidence: number | null; + probabilities: Record; +} + +export type ClassificationAnswer = BooleanAnswer | ChoiceAnswer | ScoreAnswer; + +export interface ClassificationUsage { + inputTokens: number; + outputTokens: number; +} + +export interface ClassificationRequest { + state: ClassificationState; + questions: Record; + signal?: AbortSignal; + label?: string; + /** Overrides the provider's timeout for this request alone. */ + timeoutMs?: number; +} + +export interface ClassificationResult { + model: string; + answers: Record; + usage: ClassificationUsage; +} + +export interface Classifier { + readonly id: string; + readonly model: string; + classify(request: ClassificationRequest): Promise; +} + +export type ClassificationFailure = + | 'timeout' + | 'aborted' + | 'rate_limited' + | 'unauthorized' + | 'bad_request' + | 'server_error' + | 'network' + | 'unsupported_question' + | 'malformed_response'; + +export class ClassificationError extends Error { + readonly failure: ClassificationFailure; + readonly provider: string; + readonly status?: number; + retryAfterMs?: number; + + constructor( + failure: ClassificationFailure, + message: string, + options?: { provider?: string; status?: number }, + ) { + super(message); + this.name = 'ClassificationError'; + this.failure = failure; + this.provider = options?.provider ?? 'unknown'; + this.status = options?.status; + } +} + +export function isBooleanAnswer(answer: ClassificationAnswer | undefined): answer is BooleanAnswer { + return answer?.type === 'boolean'; +} + +export function isChoiceAnswer(answer: ClassificationAnswer | undefined): answer is ChoiceAnswer { + return answer?.type === 'choice'; +} + +export function isScoreAnswer(answer: ClassificationAnswer | undefined): answer is ScoreAnswer { + return answer?.type === 'score'; +} diff --git a/packages/api/src/index.ts b/packages/api/src/index.ts index 83a0cc9b63d..0113e5a5e44 100644 --- a/packages/api/src/index.ts +++ b/packages/api/src/index.ts @@ -97,6 +97,8 @@ export * from './images'; export * from './storage'; /* Tools */ export * from './tools'; +/* Classification */ +export * from './classification'; /* web search */ export * from './web'; /* Langfuse */ diff --git a/packages/data-provider/src/config.ts b/packages/data-provider/src/config.ts index 969b3305134..9d7645289a7 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -2908,12 +2908,46 @@ export type TOpenIdDiscoveryConfig = z.infer; /** Maximum CAS attempts per ACL document, including the initial attempt. */ export const permissionWriteAttemptsSchema = z.number().int().min(1).max(100).default(3); +export const classificationProviderSchema = z + .object({ + baseURL: z.string().url().optional(), + model: z.string().optional(), + /** Which wire vocabulary the endpoint speaks. */ + dialect: z.enum(['port', 'systemone']).optional(), + /** Nests `state` and `questions` under this key, for hosts that wrap them. */ + requestKey: z.string().optional(), + /** Reads the answer envelope from this key, for hosts that wrap the response. */ + responseKey: z.string().optional(), + /** Per-request ceiling. A judgment that misses it is abandoned, never awaited. */ + timeoutMs: z.number().int().positive().max(60_000).optional(), + /** Retries for a rate limit or a server error only. */ + maxRetries: z.number().int().nonnegative().max(5).optional(), + /** Environment variable holding this provider's key. Never the key itself. */ + apiKeyEnv: z.string().optional(), + }) + /** Strict so a misspelled key fails loudly here rather than as a missing + * setting much later. */ + .strict(); + +export type TClassificationProviderConfig = z.infer; + +/** Every capability defaults to off, so an unset block changes nothing. */ +export const classificationSchema = z.object({ + enabled: z.boolean().default(false), + /** Which registered provider answers. An unknown name disables classification. */ + provider: z.string().default('http'), + providers: z.record(z.string(), classificationProviderSchema).default({}), +}); + +export type TClassificationConfig = z.infer; + export const configSchema = z.object({ version: z.string(), permissions: z.object({ maxWriteAttempts: permissionWriteAttemptsSchema }).optional(), cache: z.boolean().default(true), ocr: ocrSchema.optional(), webSearch: webSearchSchema.optional(), + classification: classificationSchema.optional(), langfuse: langfuseConfigSchema.optional(), memory: memorySchema.optional(), summarization: summarizationConfigSchema.optional(), diff --git a/packages/data-schemas/src/app/service.ts b/packages/data-schemas/src/app/service.ts index e762d1620e8..021c8763f55 100644 --- a/packages/data-schemas/src/app/service.ts +++ b/packages/data-schemas/src/app/service.ts @@ -161,6 +161,7 @@ export const AppService = async (params?: { const mcpServersConfig = config.mcpServers || null; const mcpSettings = config.mcpSettings || null; + const classification = config.classification || null; const actions = config.actions; const registration = config.registration ?? configDefaults.registration; const interfaceConfig = await loadDefaultInterface({ config, configDefaults }); @@ -181,6 +182,7 @@ export const AppService = async (params?: { skillSync, webSearch, mcpSettings, + classification, fileStrategy, registration, transactions, diff --git a/packages/data-schemas/src/types/app.ts b/packages/data-schemas/src/types/app.ts index df2a4b60540..fc188edb50e 100644 --- a/packages/data-schemas/src/types/app.ts +++ b/packages/data-schemas/src/types/app.ts @@ -102,6 +102,8 @@ export interface AppConfig { mcpConfig?: TCustomConfig['mcpServers'] | null; /** MCP settings (domain allowlist, etc.) */ mcpSettings?: TCustomConfig['mcpSettings'] | null; + /** Classification provider and the capabilities that consult it */ + classification?: TCustomConfig['classification'] | null; /** File configuration */ fileConfig?: TFileConfig; /** Secure image links configuration, enabled unless explicitly disabled */