Skip to content
Merged
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
35 changes: 29 additions & 6 deletions src/graphs/Graph.ts
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,10 @@ import {
InvalidModelToolCallError,
snapshotAcceptedModelResponse,
} from './acceptedModelResponse';
import {
collectSubagentHostArgNames,
pickSubagentHostArgInput,
} from '@/tools/subagent/hostArgs';
import {
prepareProviderRequest,
usesNativeOpenAIResponses,
Expand Down Expand Up @@ -1436,7 +1440,11 @@ export class StandardGraph extends Graph<t.BaseGraphState, t.GraphNode> {
* paths must not run in that state. */
protected resolveTrippedBreakerReason(
breakerSignal: AbortSignal = this.breakerAbort.signal
): StreamLimitExceededError | PreparedSubagentError | ProviderTextProtectionError | undefined {
):
| StreamLimitExceededError
| PreparedSubagentError
| ProviderTextProtectionError
| undefined {
if (
breakerSignal.aborted &&
(breakerSignal.reason instanceof StreamLimitExceededError ||
Expand Down Expand Up @@ -1593,7 +1601,10 @@ export class StandardGraph extends Graph<t.BaseGraphState, t.GraphNode> {
validateProviderTextProtection(providerTextProtection);
this.providerTextProtection = providerTextProtection;
this.toolExecution = toolExecution;
if (clientDelegatedToolNames != null && clientDelegatedToolNames.length > 0) {
if (
clientDelegatedToolNames != null &&
clientDelegatedToolNames.length > 0
) {
if (agents.length !== 1) {
throw new Error('Client tool delegation requires a single-agent graph');
}
Expand Down Expand Up @@ -4420,7 +4431,8 @@ export class StandardGraph extends Graph<t.BaseGraphState, t.GraphNode> {
primaryError instanceof InvalidModelToolCallError ||
primaryError instanceof ProviderTextProtectionError
) {
if (primaryError instanceof ProviderTextProtectionError) attemptBreaker.abort(primaryError);
if (primaryError instanceof ProviderTextProtectionError)
attemptBreaker.abort(primaryError);
throw primaryError;
}
if (
Expand Down Expand Up @@ -4767,7 +4779,8 @@ export class StandardGraph extends Graph<t.BaseGraphState, t.GraphNode> {
fallbackError instanceof InvalidModelToolCallError ||
fallbackError instanceof ProviderTextProtectionError
) {
if (fallbackError instanceof ProviderTextProtectionError) attemptBreaker.abort(fallbackError);
if (fallbackError instanceof ProviderTextProtectionError)
attemptBreaker.abort(fallbackError);
throw fallbackError;
}
if (
Expand Down Expand Up @@ -5382,6 +5395,7 @@ export class StandardGraph extends Graph<t.BaseGraphState, t.GraphNode> {
): GraphFactory => snapshotChildGraphFactory(parentHandlerRegistry),
});
this.registerSubagentExecutor(executor);
const hostArgNames = collectSubagentHostArgNames(executableConfigs);

const subagentTool = tool(
async (rawInput, config) => {
Expand All @@ -5391,6 +5405,10 @@ export class StandardGraph extends Graph<t.BaseGraphState, t.GraphNode> {
subagent_thread_id?: string;
run_in_background?: boolean;
};
const hostArgInput = pickSubagentHostArgInput(input, hostArgNames);
if (!hostArgInput.ok) {
return hostArgInput.message;
}
const description =
typeof input.description === 'string' &&
input.description.trim().length > 0
Expand Down Expand Up @@ -5447,6 +5465,9 @@ export class StandardGraph extends Graph<t.BaseGraphState, t.GraphNode> {
parentConfigurable: config.configurable as
| Record<string, unknown>
| undefined,
...(hostArgInput.hostArgs == null
? {}
: { hostArgs: hostArgInput.hostArgs }),
};
if (input.run_in_background === true) {
return executor.executeInBackground({
Expand Down Expand Up @@ -5577,8 +5598,10 @@ export class StandardGraph extends Graph<t.BaseGraphState, t.GraphNode> {
const delegatedNames = this.clientDelegatedToolNames;
if (delegatedNames != null && delegatedNames.size > 0) {
const { messages } = state as t.BaseGraphState;
const last = messages[messages.length - 1] as AIMessageChunk | undefined;
const calls = last?.getType() === 'ai' ? last.tool_calls ?? [] : [];
const last = messages[messages.length - 1] as
| AIMessageChunk
| undefined;
const calls = last?.getType() === 'ai' ? (last.tool_calls ?? []) : [];
if (calls.some((call) => delegatedNames.has(call.name))) {
if (
calls.some(
Expand Down
174 changes: 174 additions & 0 deletions src/specs/subagent-host-args.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,174 @@
import { HumanMessage } from '@langchain/core/messages';
import { FakeListChatModel } from '@langchain/core/utils/testing';
import type { RunnableConfig } from '@langchain/core/runnables';
import type { ToolCall } from '@langchain/core/messages/tool';
import type * as t from '@/types';
import { Constants, GraphEvents, Providers, ToolEndHandler } from '@/index';
import * as providers from '@/llm/providers';
import { Run } from '@/run';

const CHILD_RESPONSE = 'Reviewed the pull request on the selected machine.';

const callerConfig: Partial<RunnableConfig> & {
version: 'v1' | 'v2';
streamMode: string;
} = {
configurable: { thread_id: 'subagent-host-args-thread' },
streamMode: 'values',
version: 'v2' as const,
};

const childInputs = (agentId: string): t.AgentInputs => ({
agentId,
provider: Providers.OPENAI,
clientOptions: { modelName: 'gpt-4o-mini', apiKey: 'test-key' },
instructions: `You are ${agentId}.`,
maxContextTokens: 8000,
});

const createParentAgent = (
resolveReviewer: jest.Mock<Promise<t.AgentInputs>, [t.SubagentResolveContext]>
): t.AgentInputs => ({
agentId: 'parent',
provider: Providers.OPENAI,
clientOptions: { modelName: 'gpt-4o-mini', apiKey: 'test-key' },
instructions: 'Delegate reviews with the subagent tool.',
maxContextTokens: 8000,
subagentConfigs: [
{
type: 'reviewer',
name: 'PR Reviewer',
description: 'Reviews pull requests.',
configId: 'reviewer@v1',
resolveAgentInputs: resolveReviewer,
hostArgs: {
machine: {
description: 'Code machine the reviewer runs on.',
enum: ['byom-laptop', 'code-api'],
},
},
},
{
type: 'researcher',
name: 'Researcher',
description: 'Researches topics.',
agentInputs: childInputs('researcher'),
},
],
});

async function runWithToolCall(
resolveReviewer: jest.Mock<
Promise<t.AgentInputs>,
[t.SubagentResolveContext]
>,
args: Record<string, string>
): Promise<string> {
const run = await Run.create<t.IState>({
runId: `subagent-host-args-${Date.now()}`,
graphConfig: {
type: 'standard',
agents: [createParentAgent(resolveReviewer)],
},
returnContent: true,
skipCleanup: true,
customHandlers: { [GraphEvents.TOOL_END]: new ToolEndHandler() },
});
const toolCall: ToolCall = {
id: 'call_review',
name: Constants.SUBAGENT,
args,
type: 'tool_call',
};
run.Graph?.overrideTestModel(['Delegating.', 'Done.'], 10, [toolCall]);
await run.processStream(
{ messages: [new HumanMessage('Review the PR.')] },
callerConfig
);
const result = run
.getRunMessages()
?.find(
(message) =>
message._getType() === 'tool' &&
'name' in message &&
message.name === Constants.SUBAGENT
);
return String(result?.content ?? '');
}

describe('subagent host arguments through a run', () => {
jest.setTimeout(30000);

let getChatModelClassSpy: jest.SpyInstance;
const originalGetChatModelClass = providers.getChatModelClass;

beforeEach(() => {
getChatModelClassSpy = jest
.spyOn(providers, 'getChatModelClass')
.mockImplementation(((provider: Providers) => {
if (provider === Providers.OPENAI) {
return class extends FakeListChatModel {
// eslint-disable-next-line @typescript-eslint/no-explicit-any
constructor(_options: any) {
super({ responses: [CHILD_RESPONSE] });
}
// eslint-disable-next-line @typescript-eslint/no-explicit-any
} as any;
}
return originalGetChatModelClass(provider);
}) as typeof providers.getChatModelClass);
});

afterEach(() => {
getChatModelClassSpy.mockRestore();
});

it('delivers the parent’s choice to the selected resolver', async () => {
const resolveReviewer = jest.fn(
async (_context: t.SubagentResolveContext) => childInputs('reviewer')
);

const content = await runWithToolCall(resolveReviewer, {
description: 'Review PR #600.',
subagent_type: 'reviewer',
machine: 'byom-laptop',
});

expect(content).toContain(CHILD_RESPONSE);
expect(resolveReviewer).toHaveBeenCalledTimes(1);
expect(resolveReviewer.mock.calls[0][0].hostArgs).toEqual({
machine: 'byom-laptop',
});
});

it('falls back to the host default when the parent omits the argument', async () => {
const resolveReviewer = jest.fn(
async (_context: t.SubagentResolveContext) => childInputs('reviewer')
);

const content = await runWithToolCall(resolveReviewer, {
description: 'Review PR #600.',
subagent_type: 'reviewer',
});

expect(content).toContain(CHILD_RESPONSE);
expect(resolveReviewer.mock.calls[0][0].hostArgs).toBeUndefined();
});

it('returns a model-visible error for an argument the selected type does not accept', async () => {
const resolveReviewer = jest.fn(
async (_context: t.SubagentResolveContext) => childInputs('reviewer')
);

const content = await runWithToolCall(resolveReviewer, {
description: 'Research the topic.',
subagent_type: 'researcher',
machine: 'code-api',
});

expect(content).toContain(
'Subagent "researcher" does not accept "machine". Omit it for this subagent type.'
);
expect(resolveReviewer).not.toHaveBeenCalled();
});
});
17 changes: 14 additions & 3 deletions src/tools/SubagentTool.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import type { JsonSchemaType, LCTool } from '@/types/tools';
import type { SubagentConfig } from '@/types';
import { buildSubagentHostArgProperties } from '@/tools/subagent/hostArgs';
import { INTENT_PROPERTY } from '@/tools/intentArg';
import { Constants } from '@/common';

Expand Down Expand Up @@ -33,6 +34,9 @@ const RUN_IN_BACKGROUND_PROP_DESCRIPTION =
const SUBAGENT_THREAD_PROP_DESCRIPTION =
'Continue a host-owned child thread using a fresh execution lease. The saved thread must belong to this scope and subagent type. Only available with run_in_background.';

const HOST_ARGS_DESCRIPTION =
'\n\nOPTIONAL ARGUMENTS:\n- Some types accept extra arguments, listed in brackets after the type. Each is optional: omit it to let the host choose. Where values are listed, pass one listed for the selected type; where it says any text, pass a short value of your own.';

export const SubagentToolSchema = {
type: 'object',
properties: {
Expand Down Expand Up @@ -69,8 +73,13 @@ export function buildSubagentToolParams(
description: string;
} {
const types = configs.map((c) => c.type);
const hostArgs = buildSubagentHostArgProperties(configs);
const typeDescriptions = configs
.map((c) => `- "${c.type}" (${c.name}): ${c.description}`)
.map((c) => {
const summary = hostArgs.summaries.get(c.type);
const line = `- "${c.type}" (${c.name}): ${c.description}`;
return summary == null ? line : `${line} [optional ${summary}]`;
})
.join('\n');

return {
Expand All @@ -96,19 +105,21 @@ export function buildSubagentToolParams(
},
}
: {}),
...(options.background === true &&
options.threadContinuation === true
...(options.background === true && options.threadContinuation === true
? {
subagent_thread_id: {
type: 'string',
description: SUBAGENT_THREAD_PROP_DESCRIPTION,
},
}
: {}),
...hostArgs.properties,
},
required: ['description', 'subagent_type'],
},
description: `${SubagentToolDescription}${
hostArgs.summaries.size > 0 ? HOST_ARGS_DESCRIPTION : ''
}${
options.background === true
? '\n\nBACKGROUND EXECUTION:\n- Set run_in_background to true when you do not need the result immediately. The call returns a background_task_id; use the host background-task tools to poll, steer, queue, interrupt, or cancel it.'
: ''
Expand Down
8 changes: 6 additions & 2 deletions src/tools/subagent/SubagentExecutionRegistry.ts
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,8 @@ export type SubagentInvocationBinding = Readonly<{
description: string;
subagentType: string;
configId?: string;
/** Value-free identity of the call's validated host arguments. */
hostArgsDigest?: string;
}>;

export type SubagentSettlementBinding = Readonly<{
Expand Down Expand Up @@ -340,7 +342,8 @@ function assertCompatibleResumeExecution(
current.parentToolCallId === next.parentToolCallId &&
current.childRunId === next.childRunId &&
current.subagentType === next.subagentType &&
current.configId === next.configId
current.configId === next.configId &&
current.hostArgsDigest === next.hostArgsDigest
) {
return;
}
Expand All @@ -354,7 +357,8 @@ function assertSameInvocation(
if (
current.description === next.description &&
current.subagentType === next.subagentType &&
current.configId === next.configId
current.configId === next.configId &&
current.hostArgsDigest === next.hostArgsDigest
) {
return;
}
Expand Down
Loading
Loading