Skip to content
Closed
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
9 changes: 9 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -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.<name>.apiKeyEnv`; this is the default.
# CLASSIFIER_API_KEY=your_classifier_api_key

#======================#
# MCP Configuration #
#======================#
Expand Down
41 changes: 41 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 @@ -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;
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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 =
Expand Down Expand Up @@ -4916,6 +4956,7 @@ class AgentClient extends BaseClient {
messages,
discoveredToolNames:
this.eventActorContinuation === 'warm' ? this.eventActorDiscoveredToolNames : undefined,
predictedToolNames,
modelCallbacks: [
modelBoundCallback,
createAgentMemoryCallback(this.attachmentMemoryContext ?? {}),
Expand Down
41 changes: 41 additions & 0 deletions librechat.example.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -515,6 +515,47 @@ 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
#
# # 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.
# # 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
21 changes: 20 additions & 1 deletion packages/api/src/agents/memory.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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';
Expand Down Expand Up @@ -1031,6 +1032,7 @@ export async function createMemoryProcessor({
jobCreatedAt,
user,
tenantId,
gate,
}: {
res: ServerResponse;
messageId: string;
Expand All @@ -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,
Expand Down Expand Up @@ -1074,6 +1078,21 @@ export async function createMemoryProcessor({
messages: BaseMessage[],
inspectionMessages?: BaseMessage[],
): Promise<(TAttachment | null)[] | 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({
res,
Expand All @@ -1093,7 +1112,7 @@ export async function createMemoryProcessor({
totalTokens: totalTokens || 0,
tokenCountsByKey,
filters,
instructions: finalInstructions,
instructions: turnInstructions,
setMemory: memoryMethods.setMemory,
deleteMemory: memoryMethods.deleteMemory,
user,
Expand Down
11 changes: 11 additions & 0 deletions packages/api/src/agents/run.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2074,6 +2074,7 @@ export async function createRun({
agents,
messages,
discoveredToolNames,
predictedToolNames,
requestBody,
codeApprovalMode: requestedCodeApprovalMode,
user,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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}. */
Expand Down
5 changes: 5 additions & 0 deletions packages/api/src/classification/index.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
export * from './types';
export * from './questions';
export * from './registry';
export * from './resolve';
export type { ProviderFetch, Transport, TransportOptions } from './providers/transport';
Loading