Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -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.<name>.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 #
#======================#
Expand Down
70 changes: 70 additions & 0 deletions api/server/controllers/agents/client.js
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,9 @@ const {
computeAgentRequestFingerprint,
computeLegacyAgentRequestFingerprint,
getRunDiscoveredTools,
predictToolsForTurn,
createMemoryGate,
classificationCapability,
captureResumeModelParameters,
pickResumeContext,
getApprovalTtlMs,
Expand Down Expand Up @@ -1852,6 +1855,40 @@ class AgentClient extends BaseClient {
return wiring;
}

/**
* Builds the classifier gate that skips the memory model on turns with nothing durable.
* @returns {import('@librechat/api').MemoryGate | null}
*/
buildMemoryGate() {
const config = this.options.req?.config?.classification;
const capability = classificationCapability(config, 'memoryGate');
if (capability == null) {
return null;
}
return createMemoryGate({
classifier: capability.classifier,
settings: capability.settings,
onUsage: (usage, model) => this.recordClassifierUsage(usage, model),
});
}

/** Bills a classifier judgment like any other secondary call: priced from the shared table. */
recordClassifierUsage(usage, model) {
const appConfig = this.options.req?.config;
return this.recordCollectedUsage({
collectedUsage: [{ input_tokens: usage.inputTokens, output_tokens: usage.outputTokens }],
context: 'classification',
model,
crossEndpoint: true,
balance: getBalanceConfig(appConfig),
transactions: getTransactionsConfig(appConfig),
messageId: this.responseMessageId,
updateStreamUsage: false,
}).catch((err) => {
logger.error('[AgentClient] Error recording classifier usage', getSafeErrorMetadata(err));
});
}

/** Builds the independently opt-in live reasoning-label controller. */
buildReasoningLabelWiring(streamId, abortSignal, seedFromContent = false) {
if (!streamId || typeof Run?.prototype?.generateReasoningLabel !== 'function') {
Expand Down Expand Up @@ -3383,6 +3420,8 @@ class AgentClient extends BaseClient {
res: this.options.res,
user: createSafeUser(this.options.req.user),
tenantId: resolveRequestTenantId(this.options.req),
/** Null unless `classification.memoryGate` is on. */
gate: this.buildMemoryGate(),
});

this.processMemory = processMemory;
Expand Down Expand Up @@ -4378,6 +4417,16 @@ class AgentClient extends BaseClient {
);
}

/** By the pause these describe what the turn had loaded, so the resumed
* segment must rebuild with them. Run state records real discoveries only. */
if (this.predictedToolNames?.length) {
const merged = new Set(discoveredTools);
for (const name of this.predictedToolNames) {
merged.add(name);
}
discoveredTools = Array.from(merged);
}

this.stagedApproval = {
streamId,
pendingAction,
Expand Down Expand Up @@ -4798,6 +4847,21 @@ class AgentClient extends BaseClient {
if (this.agentConfigs && this.agentConfigs.size > 0) {
agents.push(...this.agentConfigs.values());
}

/** Awaited after the memory run starts, so the two overlap. */
const predictionPromise = predictToolsForTurn({
config: appConfig?.classification,
agents,
messages,
signal: abortController.signal,
onUsage: (usage, model) => this.recordClassifierUsage(usage, model),
}).catch((error) => {
logger.warn(
'[AgentClient] Tool prediction failed; continuing without it',
getSafeErrorMetadata(error),
);
return [];
});
const modelBoundCallback =
AgentClient.prototype.createModelBoundChatModelCallback.call(this);
const initialModelBoundAdmission =
Expand Down Expand Up @@ -4840,6 +4904,11 @@ class AgentClient extends BaseClient {
memoryPromise = this.runMemory(memoryMessages);
}

const predictedToolNames = await predictionPromise;
if (predictedToolNames.length > 0) {
this.predictedToolNames = predictedToolNames;
}

const { calibrationRatio, fadingTier, fadingTiers } = resolveRunSeeds(this);

const streamId = this.options.req?._resumableStreamId;
Expand Down Expand Up @@ -4916,6 +4985,7 @@ class AgentClient extends BaseClient {
messages,
discoveredToolNames:
this.eventActorContinuation === 'warm' ? this.eventActorDiscoveredToolNames : undefined,
predictedToolNames,
modelCallbacks: [
modelBoundCallback,
createAgentMemoryCallback(this.attachmentMemoryContext ?? {}),
Expand Down
93 changes: 93 additions & 0 deletions api/server/controllers/agents/client.test.js
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,9 @@ const mockDetachedUsageRecorder = jest.fn();
const mockCreateDetachedSubagentUsageRecorder = jest.fn(() => mockDetachedUsageRecorder);
const mockGetAgentCheckpointer = jest.fn();
const mockHasDurableAgentInterruptCheckpoint = jest.fn().mockResolvedValue(true);
const mockPredictToolsForTurn = jest.fn((...args) =>
jest.requireActual('@librechat/api').predictToolsForTurn(...args),
);
const mockBuildAgentScopedContext = jest.fn((...args) =>
jest.requireActual('@librechat/api').buildAgentScopedContext(...args),
);
Expand Down Expand Up @@ -190,6 +193,7 @@ jest.mock('@librechat/api', () => ({
buildAgentScopedContext: (...args) => mockBuildAgentScopedContext(...args),
checkAccess: jest.fn(),
createRun: (...args) => mockCreateRun(...args),
predictToolsForTurn: (...args) => mockPredictToolsForTurn(...args),
countFormattedMessageTokens: jest.fn(() => 42),
countTokens: jest.fn((text) => Math.ceil(String(text ?? '').length / 4)),
createCachedTokenCounter: jest.fn(async () => jest.fn(() => 0)),
Expand Down Expand Up @@ -2377,6 +2381,95 @@ describe('AgentClient - startup telemetry', () => {
},
);

it('bills classifier usage as its own secondary transaction', async () => {
mockRecordCollectedUsage.mockClear();
const client = new AgentClient({
req: { user: { id: 'user-123' }, body: {}, config: {} },
res: {},
agent: {
id: 'agent-123',
endpoint: EModelEndpoint.openAI,
provider: EModelEndpoint.openAI,
model_parameters: { model: 'gpt-4' },
},
endpointTokenConfig: {},
eventHandlers: {},
contentParts: [],
collectedUsage: [],
artifactPromises: [],
});
client.responseMessageId = 'classifier-usage-response';
mockRecordCollectedUsage.mockResolvedValueOnce({ input_tokens: 512, output_tokens: 0 });

await client.recordClassifierUsage({ inputTokens: 512, outputTokens: 0 }, 'jev-latest');

expect(mockRecordCollectedUsage).toHaveBeenCalledTimes(1);
const [, params] = mockRecordCollectedUsage.mock.calls[0];
expect(params).toEqual(
expect.objectContaining({
context: 'classification',
model: 'jev-latest',
messageId: 'classifier-usage-response',
collectedUsage: [{ input_tokens: 512, output_tokens: 0 }],
}),
);
expect(client.getStreamUsage()).toBeFalsy();
});

it('starts the memory run while tool prediction is still running', async () => {
let resolvePrediction;
mockPredictToolsForTurn.mockImplementationOnce(
() => new Promise((resolve) => (resolvePrediction = resolve)),
);
mockCreateRun.mockResolvedValueOnce({
Graph: null,
processStream: jest.fn().mockResolvedValue(),
getCalibrationRatio: jest.fn(() => 0),
getInterrupt: jest.fn(() => undefined),
});
const createRunBefore = mockCreateRun.mock.calls.length;
const client = new AgentClient({
req: { user: { id: 'user-123' }, body: {}, config: {} },
res: {},
agent: {
id: 'agent-123',
endpoint: EModelEndpoint.openAI,
provider: EModelEndpoint.openAI,
model_parameters: { model: 'gpt-4' },
hide_sequential_outputs: false,
tools: [{ name: 'read_file' }],
},
endpointTokenConfig: {},
eventHandlers: {},
contentParts: [],
collectedUsage: [],
artifactPromises: [],
});
client.conversationId = 'prediction-overlap';
client.responseMessageId = 'prediction-overlap-response';
client.parentMessageId = 'prediction-overlap-parent';
client.recordCollectedUsage = jest.fn().mockResolvedValue();
client.processMemory = jest.fn();
client.runMemory = jest.fn().mockResolvedValue(undefined);

const completion = client.chatCompletion({ payload: [] });
for (let i = 0; i < 50 && client.runMemory.mock.calls.length === 0; i++) {
await new Promise((resolve) => setImmediate(resolve));
}

expect(resolvePrediction).toBeDefined();
expect(client.runMemory).toHaveBeenCalledTimes(1);
expect(mockCreateRun).toHaveBeenCalledTimes(createRunBefore);

resolvePrediction(['search_mcp_docs']);
await completion;

expect(mockCreateRun).toHaveBeenCalledTimes(createRunBefore + 1);
expect(mockCreateRun.mock.calls[createRunBefore][0].predictedToolNames).toEqual([
'search_mcp_docs',
]);
});

it('uses request-scoped hook resolution when deciding whether a scheduled run can pause', async () => {
mockIsHITLEnabled.mockReturnValue(true);
registerToolApprovalHook((context) =>
Expand Down
71 changes: 71 additions & 0 deletions librechat.example.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -515,6 +515,77 @@ actions:
# - 'host.docker.internal:8080'
# - '127.0.0.1:8080'

# Classification: small typed judgments (yes/no, pick one) that code can branch
# on. Off unless enabled. The key is read from the environment variable named by
# apiKeyEnv, never written here. `timeoutMs` bounds each judgment, retries
# included. Every enabled capability sends conversation text to the configured
# host: tool selection sends the latest request and the deferred tools'
# descriptions, the memory gate the last few messages.
# classification:
# enabled: true
# provider: http
# providers:
# http:
# baseURL: https://classifier.example.com/v1/classify
# apiKeyEnv: CLASSIFIER_API_KEY
# timeoutMs: 4000
#
# # Known hosts ship as presets, so naming one is usually enough. Cloudflare
# # is the exception: its URL carries your account id.
# provider: cloudflare
# providers:
# cloudflare:
# baseURL: https://api.cloudflare.com/client/v4/accounts/<account-id>/ai/run
# apiKeyEnv: CLOUDFLARE_API_TOKEN
#
# # A host with no preset needs no code either, only its shape. `dialect`
# # picks the wire vocabulary, `requestKey` nests the body, `responseKey`
# # unwraps the reply.
# provider: inhouse
# providers:
# inhouse:
# baseURL: https://classify.internal/v1/run
# model: your-model
# dialect: systemone
# requestKey: input
# responseKey: result
# apiKeyEnv: INHOUSE_CLASSIFIER_KEY
#
# # Surfaces the deferred tools a turn is likely to need, so their schemas ship
# # with the first model call instead of costing a tool_search round trip.
# # Only ever adds: a tool it passes over stays listed by name and one search away.
# toolSelection:
# enabled: true
# shortlist: 5
# # Overrides the provider timeout. It runs before the first model call, so
# # it adds up to this much latency to a turn; a large catalog may need more.
# timeoutMs: 6000
# # Below this probability that any tool is needed, only tools the request
# # names outright are surfaced.
# needsToolThreshold: 0.15
# # An unsure ranking surfaces this many extra rather than fewer.
# lowConfidenceExtra: 3
# # Replace the wording of either question without touching code.
# instructions: Which tool should the assistant call first?
# guidance: Prefer the tool whose purpose matches the request.
#
# # Skips the memory model on turns that do not ask to remember, update or
# # forget anything, which is all the default memory instructions act on.
# memoryGate:
# enabled: true
# threshold: 0.5
# # If your memory.instructions also keep facts the user did not ask to
# # keep, ask about those instead. whenTrue/whenFalse rather than
# # true/false: YAML reads those bare keys as booleans.
# instructions: Does `latest` say something about the user worth keeping?
# whenTrue: A lasting preference or a fact about the user.
# whenFalse: Small talk, or a detail that only matters in this task.
# # Also ask which of memory.validKeys the turn belongs under and suggest it
# # to the memory model. Rides in the request the gate already makes.
# categorize: true
# categoryThreshold: 0.4
# detectUpdates: true

# Example MCP Servers Object Structure
# mcpServers:
# everything:
Expand Down
Loading