Skip to content
Merged
275 changes: 275 additions & 0 deletions api/server/controllers/agents/__tests__/callbacks.spec.js
Original file line number Diff line number Diff line change
Expand Up @@ -25,12 +25,15 @@ jest.mock('@librechat/api', () => ({
isCodeArtifactToolOutput: jest.requireActual('@librechat/api').isCodeArtifactToolOutput,
isCodeSessionToolName: jest.requireActual('@librechat/api').isCodeSessionToolName,
collectToolCallIds: jest.requireActual('@librechat/api').collectToolCallIds,
captureSubagentIdentity: jest.requireActual('@librechat/api').captureSubagentIdentity,
createToolTimingAdapter: jest.requireActual('@librechat/api').createToolTimingAdapter,
}));

jest.mock('@librechat/data-schemas', () => ({
logger: {
debug: jest.fn(),
error: jest.fn(),
warn: jest.fn(),
},
}));

Expand Down Expand Up @@ -361,6 +364,139 @@ describe('resumable event generation fencing', () => {
expect(resumedPublish.mock.calls[0][0].activityEventId).not.toBe(firstUpdate.activityEventId);
});

it('publishes tool preparation and handoff into event-child activity', async () => {
const { GraphEvents, createContentAggregator } = jest.requireActual('@librechat/agents');
const { getDefaultHandlers } = require('../callbacks');
const publish = jest.fn().mockResolvedValue(undefined);
const { contentParts, stepMap, aggregateContent } = createContentAggregator();
const handlers = getDefaultHandlers({
res: { write: jest.fn() },
aggregateContent,
contentParts,
stepMap,
toolEndCallback: jest.fn(),
collectedUsage: [],
streamId: 'event-thread',
eventChildActivity: {
runId: 'event-thread',
parentRunId: 'parent-conversation',
subagentRunId: 'child-1',
subagentType: 'researcher',
subagentAgentId: 'agent-1',
parentAgentId: 'director',
publish,
},
});
const step = {
id: 'step-child',
index: 0,
type: 'tool_calls',
stepDetails: {
type: 'tool_calls',
tool_calls: [{ id: 'call-child', name: 'query', args: '{}' }],
},
};
await handlers[GraphEvents.ON_RUN_STEP].handle(GraphEvents.ON_RUN_STEP, step);
await handlers[GraphEvents.ON_RUN_STEP_DELTA].handle(GraphEvents.ON_RUN_STEP_DELTA, {
id: 'step-child',
observed_at: 100,
delta: { type: 'tool_calls', tool_calls: [{ id: 'call-child', index: 0, args: '{' }] },
});
await handlers[StepEvents.ON_TOOL_CALLS_DISPATCHED].handle(
StepEvents.ON_TOOL_CALLS_DISPATCHED,
{
dispatched_at: 500,
toolCalls: [{ id: 'call-child', name: 'query', stepId: 'step-child' }],
},
);
await new Promise((resolve) => setImmediate(resolve));
expect(publish.mock.calls.map(([value]) => value.phase)).toEqual([
'run_step',
'tool_preparation',
'run_step_delta',
'tool_calls_dispatched',
]);
expect(publish.mock.calls[1][0].data).toEqual({
id: 'step-child',
index: 0,
toolCallId: 'call-child',
observed_at: 100,
});
expect(publish.mock.calls[3][0].data.toolCalls[0]).not.toHaveProperty('args');
});

it('folds child dispatch and result into the parent-owned subagent tool part', async () => {
const { GraphEvents } = jest.requireActual('@librechat/agents');
const { getDefaultHandlers } = require('../callbacks');
const aggregators = new Map();
const handlers = getDefaultHandlers({
res: { write: jest.fn() },
aggregateContent: jest.fn(),
toolEndCallback: jest.fn(),
collectedUsage: [],
subagentAggregatorsByToolCallId: aggregators,
});
const base = {
parentToolCallId: 'parent-call',
parentRunId: 'parent-run',
subagentRunId: 'child-run',
subagentType: 'researcher',
subagentAgentId: 'child-agent',
runId: 'parent-run',
};
for (const event of [
{
phase: 'run_step',
data: {
id: 'child-step',
index: 0,
type: 'tool_calls',
stepDetails: {
type: 'tool_calls',
tool_calls: [{ id: 'child-call', name: 'query', args: '{}' }],
},
},
},
{
phase: 'tool_preparation',
data: { id: 'child-step', toolCallId: 'child-call', observed_at: 100 },
},
{
phase: 'tool_calls_dispatched',
data: {
dispatched_at: 500,
toolCalls: [{ id: 'child-call', stepId: 'child-step', name: 'query' }],
},
},
{
phase: 'run_step_completed',
data: {
result: {
id: 'child-step',
index: 0,
type: 'tool_call',
completed_at: 540,
tool_call: { id: 'child-call', name: 'query', args: '{}', output: 'ok', progress: 1 },
},
},
},
]) {
await handlers[GraphEvents.ON_SUBAGENT_UPDATE].handle(GraphEvents.ON_SUBAGENT_UPDATE, {
...base,
...event,
});
}
expect(jest.requireMock('@librechat/data-schemas').logger.warn).not.toHaveBeenCalled();
expect(aggregators.get('parent-call')?.contentParts[0]?.tool_call).toMatchObject({
id: 'child-call',
toolPreparationStartedAt: 100,
toolPreparationDurationMs: 400,
toolDispatchedAt: 500,
toolExecutionDurationMs: 40,
output: 'ok',
});
});

it('forwards the originating job epoch with deferred attachments', () => {
const { GenerationJobManager } = require('@librechat/api');
const { createAttachmentEmitter } = require('../callbacks');
Expand Down Expand Up @@ -1308,6 +1444,145 @@ describe('createToolEndCallback', () => {
});
});

describe('tool dispatch timing', () => {
it('forwards the SDK handoff and stores preparation and result intervals independently', async () => {
const { GraphEvents, createContentAggregator } = jest.requireActual('@librechat/agents');
const { GenerationJobManager } = require('@librechat/api');
const { getDefaultHandlers } = require('../callbacks');
const { contentParts, stepMap, aggregateContent } = createContentAggregator();
const handlers = getDefaultHandlers({
res: { write: jest.fn() },
contentParts,
stepMap,
aggregateContent,
toolEndCallback: jest.fn(),
collectedUsage: [],
streamId: 'run',
});
const step = {
id: 'step-1',
index: 0,
type: 'tool_calls',
stepDetails: {
type: 'tool_calls',
tool_calls: [{ id: 'call-1', name: 'query', args: '{}' }],
},
};
await handlers[GraphEvents.ON_RUN_STEP].handle(GraphEvents.ON_RUN_STEP, step);
await handlers[GraphEvents.ON_RUN_STEP_DELTA].handle(GraphEvents.ON_RUN_STEP_DELTA, {
id: 'step-1',
observed_at: 1_000,
delta: { type: 'tool_calls', tool_calls: [{ id: 'call-1', index: 0, args: '{' }] },
});
const dispatched = {
dispatched_at: 248_000,
toolCalls: [{ id: 'call-1', name: 'query', stepId: 'step-1' }],
};
await handlers[StepEvents.ON_TOOL_CALLS_DISPATCHED].handle(
StepEvents.ON_TOOL_CALLS_DISPATCHED,
dispatched,
);
await handlers[GraphEvents.ON_RUN_STEP_COMPLETED].handle(GraphEvents.ON_RUN_STEP_COMPLETED, {
result: {
id: 'step-1',
completed_at: 248_340,
index: 0,
tool_call: { id: 'call-1', name: 'query', args: '{}', output: 'ok' },
},
});
await handlers[GraphEvents.ON_RUN_STEP_CLOSED].handle(GraphEvents.ON_RUN_STEP_CLOSED, {
id: 'step-1',
index: 0,
type: 'tool_calls',
status: 'completed',
created_at: 1_000,
closed_at: 248_340,
});
expect(GenerationJobManager.emitChunk).toHaveBeenCalledWith(
'run',
{
event: StepEvents.ON_TOOL_PREPARATION,
data: {
id: 'step-1',
index: 0,
toolCallId: 'call-1',
observed_at: 1_000,
},
},
expect.anything(),
);
expect(GenerationJobManager.emitChunk).toHaveBeenCalledWith(
'run',
{ event: StepEvents.ON_TOOL_CALLS_DISPATCHED, data: dispatched },
expect.anything(),
);
expect(contentParts[0].tool_call).toMatchObject({
runStepDurationMs: 247_340,
runStepClosedAt: 248_340,
toolPreparationDurationMs: 247_000,
toolExecutionDurationMs: 340,
});
});

it('keeps preparation across a new handler created after HITL approval', async () => {
const { GraphEvents, createContentAggregator } = jest.requireActual('@librechat/agents');
const { getDefaultHandlers } = require('../callbacks');
const { contentParts, stepMap, aggregateContent } = createContentAggregator();
const handlers = getDefaultHandlers({
res: { write: jest.fn() },
contentParts,
stepMap,
aggregateContent,
toolEndCallback: jest.fn(),
collectedUsage: [],
toolTimingReplayEvents: [
{
event: StepEvents.ON_TOOL_PREPARATION,
data: {
id: 'step-1',
index: 0,
toolCallId: 'call-1',
observed_at: 1_000,
},
},
],
});
await handlers[GraphEvents.ON_RUN_STEP].handle(GraphEvents.ON_RUN_STEP, {
id: 'step-1',
index: 0,
type: 'tool_calls',
stepDetails: {
type: 'tool_calls',
tool_calls: [{ id: 'call-1', name: 'query', args: '{}' }],
},
});
await handlers[StepEvents.ON_TOOL_CALLS_DISPATCHED].handle(
StepEvents.ON_TOOL_CALLS_DISPATCHED,
{ dispatched_at: 51_000, toolCalls: [{ id: 'call-1', name: 'query', stepId: 'step-1' }] },
);
await handlers[GraphEvents.ON_RUN_STEP_COMPLETED].handle(GraphEvents.ON_RUN_STEP_COMPLETED, {
result: {
id: 'step-1',
index: 0,
completed_at: 51_200,
tool_call: { id: 'call-1', name: 'query', args: '{}', output: 'ok' },
},
});
await handlers[GraphEvents.ON_RUN_STEP_CLOSED].handle(GraphEvents.ON_RUN_STEP_CLOSED, {
id: 'step-1',
index: 0,
type: 'tool_calls',
status: 'completed',
created_at: 1_000,
closed_at: 51_200,
});
expect(contentParts[0].tool_call).toMatchObject({
toolPreparationDurationMs: 50_000,
toolExecutionDurationMs: 200,
});
});
});

describe('tool input validation marker', () => {
it('marks the streamed result and persisted content part out of band', async () => {
const { GraphEvents, createContentAggregator } = jest.requireActual('@librechat/agents');
Expand Down
19 changes: 16 additions & 3 deletions api/server/controllers/agents/callbacks.js
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ const {
getToolInputValidationDetails,
captureSubagentIdentity,
collectToolCallIds,
createToolTimingAdapter,
} = require('@librechat/api');
const { processFileCitations } = require('~/server/services/Files/Citations');
const { processCodeOutput, runPreviewFinalize } = require('~/server/services/Files/Code/process');
Expand Down Expand Up @@ -342,10 +343,11 @@ function subagentPhaseToGraphEvent(event) {
* @param {{ aggregateContent: Function, contentParts?: Array, stepMap?: Map }} aggregator
* @param {SubagentUpdateEvent} event
*/
function feedSubagentAggregator(aggregator, event) {
function feedSubagentAggregator(aggregator, event, applyChildTiming) {
const graphEvent = subagentPhaseToGraphEvent(event);
if (graphEvent) aggregator.aggregateContent({ event: graphEvent, data: event.data });
applyChildTiming(aggregator, event);
if (!graphEvent) return;
aggregator.aggregateContent({ event: graphEvent, data: event.data });

/** The SDK aggregator intentionally projects run-step tool calls onto its
* public content shape, so host-only routing metadata is not copied. Restore
Expand Down Expand Up @@ -416,6 +418,7 @@ function getDefaultHandlers({
usageEmitSink = null,
eventChildActivity = null,
resolveMcpServerName = null,
toolTimingReplayEvents = [],
}) {
if (!res || !aggregateContent) {
throw new Error(
Expand All @@ -425,6 +428,8 @@ function getDefaultHandlers({
const eventActivityPhases = {
[GraphEvents.ON_RUN_STEP]: 'run_step',
[GraphEvents.ON_RUN_STEP_DELTA]: 'run_step_delta',
[StepEvents.ON_TOOL_PREPARATION]: 'tool_preparation',
[StepEvents.ON_TOOL_CALLS_DISPATCHED]: 'tool_calls_dispatched',
[GraphEvents.ON_RUN_STEP_COMPLETED]: 'run_step_completed',
[GraphEvents.ON_RUN_STEP_CLOSED]: 'run_step_closed',
[GraphEvents.ON_MESSAGE_DELTA]: 'message_delta',
Expand Down Expand Up @@ -514,7 +519,12 @@ function getDefaultHandlers({
}
return emitForJob({ event: UsageEvents.ON_TOKEN_USAGE, data: payload });
};
const toolTiming = createToolTimingAdapter({
replayEvents: toolTimingReplayEvents,
emit: emitForJob,
});
const handlers = {
[StepEvents.ON_TOOL_CALLS_DISPATCHED]: toolTiming.dispatch,
[GraphEvents.CHAT_MODEL_END]: new ModelEndHandler(
collectedUsage,
collectedThoughtSignatures,
Expand Down Expand Up @@ -595,6 +605,7 @@ function getDefaultHandlers({
const index = stepMap?.get(stepId)?.index;
const part = typeof index === 'number' ? contentParts[index] : undefined;
if (part?.type === ContentTypes.TOOL_CALL && part.tool_call) {
toolTiming.close(part.tool_call, stepId);
part.tool_call.runStepStatus = data.status;
Object.assign(part.tool_call, getRunStepCloseMetadata(data));
/**
Expand Down Expand Up @@ -622,6 +633,7 @@ function getDefaultHandlers({
*/
handle: async (event, data, metadata) => {
aggregateContent({ event, data });
await toolTiming.delta(data);
if (data?.delta.type === StepTypes.TOOL_CALLS) {
await emitForJob({ event, data });
} else if (checkIfLastAgent(metadata?.last_agent_id, metadata?.langgraph_node)) {
Expand Down Expand Up @@ -657,6 +669,7 @@ function getDefaultHandlers({
agentId: metadata?.agent_id,
});
}
toolTiming.completed(data);
aggregateContent({ event, data });
const stepId = data?.result?.id;
const runStep = stepMap?.get(stepId);
Expand Down Expand Up @@ -769,7 +782,7 @@ function getDefaultHandlers({
}
try {
captureSubagentIdentity(aggregator, data);
feedSubagentAggregator(aggregator, data);
feedSubagentAggregator(aggregator, data, toolTiming.child);
} catch (err) {
logger.warn(
`[ON_SUBAGENT_UPDATE] Failed to aggregate phase "${data?.phase}" for tool_call ${key}: ${err?.message ?? err}`,
Expand Down
1 change: 1 addition & 0 deletions api/server/controllers/agents/resume.js
Original file line number Diff line number Diff line change
Expand Up @@ -1854,6 +1854,7 @@ const ResumeAgentController = async (req, res, next, initializeClient, addTitle)
checkpointNamespace,
foregroundRunId: mcpRequestBody.messageId,
requestBody: mcpRequestBody,
toolTimingReplayEvents: resumeState?.replayEvents,
});
client = result.client;

Expand Down
Loading
Loading