diff --git a/.changeset/soft-zebras-merge.md b/.changeset/soft-zebras-merge.md new file mode 100644 index 000000000..6c83bb273 --- /dev/null +++ b/.changeset/soft-zebras-merge.md @@ -0,0 +1,5 @@ +--- +"deepagents": patch +--- + +fix(deepagents): allow custom middleware to replace defaults by name diff --git a/libs/deepagents/src/agent.ts b/libs/deepagents/src/agent.ts index 42b73b8f4..41fa3dff1 100644 --- a/libs/deepagents/src/agent.ts +++ b/libs/deepagents/src/agent.ts @@ -1,13 +1,11 @@ import { createAgent, - createMiddleware, humanInTheLoopMiddleware, anthropicPromptCachingMiddleware, bedrockPromptCachingMiddleware, todoListMiddleware, SystemMessage, type AgentMiddleware, - type AnyAgentMiddleware, context, } from "langchain"; import type { @@ -33,6 +31,8 @@ import { StateBackend } from "./backends/state.js"; import { ConfigurationError } from "./errors.js"; import { InteropZodObject } from "@langchain/core/utils/types"; import { createCacheBreakpointMiddleware } from "./middleware/cache.js"; +import { createToolExclusionMiddleware } from "./middleware/tool_exclusion.js"; +import { mergeMiddlewareStack } from "./middleware/utils.js"; import { GENERAL_PURPOSE_SUBAGENT, type CompiledSubAgent, @@ -223,7 +223,7 @@ export function createDeepAgent< const anthropicModel = isAnthropicModel(model); const bedrockModel = isBedrockConverseModel(model); - let cacheMiddleware: AnyAgentMiddleware[] = []; + let cacheMiddleware: AgentMiddleware[] = []; if (anthropicModel) { cacheMiddleware = [ @@ -250,13 +250,15 @@ export function createDeepAgent< * Only the general-purpose subagent inherits the main agent's skills. * If a custom subagent needs skills, it must specify its own `skills` array. */ - const normalizeSubagentSpec = (input: SubAgent): SubAgent => { + const createSubagentDefaultMiddleware = ( + input: SubAgent, + ): AgentMiddleware[] => { const effectivePermissions = input.permissions ?? permissions; // Middleware for custom subagents (does NOT include skills from main agent). // Uses createSummarizationMiddleware (deepagents version) with backend support // and auto-computed defaults from model profile. - const subagentMiddleware = [ + return [ // Provides todo list management capabilities for tracking tasks. todoListMiddleware(), // Enables filesystem operations and optional long-term memory storage. @@ -274,11 +276,23 @@ export function createDeepAgent< ...(input.skills != null && input.skills.length > 0 ? [createSkillsMiddleware({ backend, sources: input.skills })] : []), - // Appends custom middleware from the subagent spec. - ...(input.middleware ?? []), - // Adds Anthropic cache controls when supported by the model. - ...cacheMiddleware, ]; + }; + + const normalizeSubagentSpec = (input: SubAgent): SubAgent => { + const subagentDefaultMiddleware = createSubagentDefaultMiddleware(input); + let subagentMiddleware = mergeMiddlewareStack( + subagentDefaultMiddleware, + input.middleware ?? [], + cacheMiddleware, + ); + + if (harnessProfile.excludedMiddleware.size > 0) { + subagentMiddleware = subagentMiddleware.filter( + (middleware) => !harnessProfile.excludedMiddleware.has(middleware.name), + ); + } + return { ...input, tools: input.tools ?? [], @@ -325,6 +339,12 @@ export function createDeepAgent< skills, tools: effectiveTools, }); + generalPurposeSpec.middleware = mergeMiddlewareStack( + generalPurposeSpec.middleware ?? [], + customMiddleware, + [], + { appendNew: false }, + ); inlineSubagents.unshift(generalPurposeSpec); } @@ -365,9 +385,8 @@ export function createDeepAgent< patchToolCallsMiddleware, ] = builtInMiddleware; - // Runtime middleware array: combine built-in + optional middleware. - // Note: The full type is handled separately via AllMiddleware. - const middleware: AnyAgentMiddleware[] = [ + // Runtime middleware array: combine core middleware, custom overrides, and tail middleware. + const coreMiddleware: AgentMiddleware[] = [ // Built-in middleware with deterministic ordering. todoMiddleware, // Optional root-level skills. @@ -380,8 +399,10 @@ export function createDeepAgent< ...(asyncSubAgents.length > 0 ? [createAsyncSubAgentMiddleware({ asyncSubAgents })] : []), - // User-provided middleware. - ...customMiddleware, + ]; + const tailMiddleware: AgentMiddleware[] = [ + // Profile middleware runs before cache middleware so it participates in prompt caching. + ...resolveMiddleware(harnessProfile.extraMiddleware), // Optional Anthropic cache controls. ...cacheMiddleware, // Optional memory support. @@ -398,44 +419,23 @@ export function createDeepAgent< ...(interruptOn ? [humanInTheLoopMiddleware({ interruptOn })] : []), ]; - // Apply profile middleware additions. Inserted before cache middleware - // so profile-injected middleware participates in prompt caching. - const profileMiddleware = resolveMiddleware(harnessProfile.extraMiddleware); - if (profileMiddleware.length > 0) { - const cacheIdx = middleware.findIndex( - (m) => m.name === "AnthropicPromptCachingMiddleware", - ); - if (cacheIdx !== -1) { - middleware.splice(cacheIdx, 0, ...profileMiddleware); - } else { - middleware.push(...profileMiddleware); - } - } + let middleware: AgentMiddleware[] = mergeMiddlewareStack( + coreMiddleware, + customMiddleware, + tailMiddleware, + ); - // Apply profile middleware exclusions. + // Apply profile middleware exclusions after custom replacement so exclusions win. if (harnessProfile.excludedMiddleware.size > 0) { const excluded = harnessProfile.excludedMiddleware; - const filtered = middleware.filter((m) => !excluded.has(m.name)); - middleware.length = 0; - middleware.push(...filtered); + middleware = middleware.filter((entry) => !excluded.has(entry.name)); } // Apply profile tool exclusions via a filtering middleware that runs // after all tool-injecting middleware. if (harnessProfile.excludedTools.size > 0) { - const excludedTools = harnessProfile.excludedTools; middleware.push( - createMiddleware({ - name: "_ToolExclusionMiddleware", - wrapModelCall: async (request: any, handler: any) => { - return handler({ - ...request, - tools: request.tools?.filter( - (t: { name: string }) => !excludedTools.has(t.name), - ), - }); - }, - }), + createToolExclusionMiddleware(harnessProfile.excludedTools), ); } diff --git a/libs/deepagents/src/middleware/subagent.test.ts b/libs/deepagents/src/middleware/subagent.test.ts index ce1b8a5d8..a2ebe85d3 100644 --- a/libs/deepagents/src/middleware/subagent.test.ts +++ b/libs/deepagents/src/middleware/subagent.test.ts @@ -28,10 +28,14 @@ import type { Serialized } from "@langchain/core/load/serializable"; import type { ChainValues } from "@langchain/core/utils/types"; import { createDeepAgent } from "../agent.js"; +import { StateBackend } from "../backends/state.js"; import { createSkillsMiddleware } from "./skills.js"; +import { createSummarizationMiddleware } from "./summarization.js"; +import { mergeMiddleware } from "./utils.js"; import { createFileData } from "../backends/utils.js"; import { createMockBackend } from "./test.js"; import { createSubAgent } from "./subagents.js"; +import { registerHarnessProfile } from "../profiles/index.js"; const createAgentMock = vi.mocked(createAgent); @@ -1229,3 +1233,250 @@ describe("createSubAgent", () => { expect(call.responseFormat).toBeUndefined(); }); }); + +describe("middleware override by name", () => { + const fakeModel = new FakeListChatModel({ responses: ["hello"] }); + + function namedMiddleware(name: string): AgentMiddleware { + return { name } as AgentMiddleware; + } + + function createCustomSummarizationMiddleware(): AgentMiddleware { + return createSummarizationMiddleware({ backend: new StateBackend() }); + } + + function getCreateAgentCall(name: string) { + const call = createAgentMock.mock.calls + .map(([params]) => params) + .find((params) => params.name === name); + if (call == null) { + throw new Error( + `Expected createAgent call for ${name}; saw ${createAgentMock.mock.calls + .map(([params]) => params.name ?? "") + .join(", ")}`, + ); + } + return call; + } + + function getMiddlewareStack(name: string): AgentMiddleware[] { + return getCreateAgentCall(name).middleware as AgentMiddleware[]; + } + + beforeEach(() => { + createAgentMock.mockClear(); + }); + + it("replaces matching middleware by name in place", () => { + const first = namedMiddleware("first"); + const original = namedMiddleware("target"); + const last = namedMiddleware("last"); + const replacement = namedMiddleware("target"); + + const merged = mergeMiddleware([first, original, last], [replacement]); + + expect(merged).toEqual([first, replacement, last]); + }); + + it("appends novel middleware after the base stack", () => { + const core = namedMiddleware("core"); + const customA = namedMiddleware("customA"); + const customB = namedMiddleware("customB"); + + const merged = mergeMiddleware([core], [customA, customB]); + + expect(merged).toEqual([core, customA, customB]); + }); + + it("uses the last same-name custom middleware as the replacement", () => { + const original = namedMiddleware("target"); + const first = namedMiddleware("target"); + const second = namedMiddleware("target"); + + const merged = mergeMiddleware([original], [first, second]); + + expect(merged).toEqual([second]); + }); + + it("replaces default main-agent middleware with same-name custom middleware", () => { + const custom = createCustomSummarizationMiddleware(); + + createDeepAgent({ model: fakeModel, name: "main", middleware: [custom] }); + + const middleware = getMiddlewareStack("main"); + const summarization = middleware.filter( + (entry) => entry.name === "SummarizationMiddleware", + ); + expect(summarization).toHaveLength(1); + expect(summarization[0]).toBe(custom); + }); + + it("keeps novel main-agent middleware before prompt caching", () => { + const anthropicModel = new FakeListChatModel({ responses: ["hello"] }); + vi.spyOn(anthropicModel, "getName").mockReturnValue("ChatAnthropic"); + const custom = namedMiddleware("CustomPromptMiddleware"); + + createDeepAgent({ + model: anthropicModel, + name: "main", + middleware: [custom], + }); + + const middleware = getMiddlewareStack("main"); + const customIndex = middleware.indexOf(custom); + const cacheIndex = middleware.findIndex( + (entry) => + entry.name === "AnthropicPromptCachingMiddleware" || + entry.name === "CacheBreakpointMiddleware", + ); + expect(customIndex).toBeGreaterThanOrEqual(0); + expect(cacheIndex).toBeGreaterThanOrEqual(0); + expect(customIndex).toBeLessThan(cacheIndex); + }); + + it("replaces prompt cache defaults in both main and general-purpose stacks", () => { + const anthropicModel = new FakeListChatModel({ responses: ["hello"] }); + vi.spyOn(anthropicModel, "getName").mockReturnValue("ChatAnthropic"); + const custom = namedMiddleware("CacheBreakpointMiddleware"); + + createDeepAgent({ + model: anthropicModel, + name: "main", + middleware: [custom], + }); + + for (const agentName of ["main", "general-purpose"]) { + const middleware = getMiddlewareStack(agentName); + const cacheEntries = middleware.filter( + (entry) => entry.name === "CacheBreakpointMiddleware", + ); + expect(cacheEntries).toHaveLength(1); + expect(cacheEntries[0]).toBe(custom); + } + }); + + it("lets profile middleware exclusions win over custom replacements", () => { + registerHarnessProfile("override-test:model", { + excludedMiddleware: ["SummarizationMiddleware"], + }); + const custom = createCustomSummarizationMiddleware(); + + createDeepAgent({ + model: "override-test:model", + name: "main", + middleware: [custom], + }); + + const middleware = getMiddlewareStack("main"); + expect( + middleware.some((entry) => entry.name === "SummarizationMiddleware"), + ).toBe(false); + }); + + it("keeps tool exclusion middleware last", () => { + registerHarnessProfile("tool-exclusion-test:model", { + excludedTools: ["write_file"], + }); + const custom = namedMiddleware("CustomToolMiddleware"); + + createDeepAgent({ + model: "tool-exclusion-test:model", + name: "main", + middleware: [custom], + }); + + const middleware = getMiddlewareStack("main"); + expect(middleware[middleware.length - 1]?.name).toBe( + "_ToolExclusionMiddleware", + ); + }); + + it("passes main-agent default overrides to the general-purpose subagent", () => { + const custom = createCustomSummarizationMiddleware(); + + createDeepAgent({ model: fakeModel, name: "main", middleware: [custom] }); + + const middleware = getMiddlewareStack("general-purpose"); + const summarization = middleware.filter( + (entry) => entry.name === "SummarizationMiddleware", + ); + expect(summarization).toHaveLength(1); + expect(summarization[0]).toBe(custom); + }); + + it("does not pass parent-only middleware to the general-purpose subagent", () => { + const custom = namedMiddleware("ParentOnlyMiddleware"); + + createDeepAgent({ model: fakeModel, name: "main", middleware: [custom] }); + + const middleware = getMiddlewareStack("general-purpose"); + expect(middleware).not.toContain(custom); + }); + + it("does not pass main-agent default overrides to declarative subagents", () => { + const custom = createCustomSummarizationMiddleware(); + + createDeepAgent({ + model: fakeModel, + name: "main", + middleware: [custom], + subagents: [ + { + name: "helper", + description: "Helps with work", + systemPrompt: "Help.", + }, + ], + }); + + const middleware = getMiddlewareStack("helper"); + expect(middleware).not.toContain(custom); + expect( + middleware.some((entry) => entry.name === "SummarizationMiddleware"), + ).toBe(true); + }); + + it("does not pass parent-only middleware to declarative subagents", () => { + const custom = namedMiddleware("ParentOnlyMiddleware"); + + createDeepAgent({ + model: fakeModel, + name: "main", + middleware: [custom], + subagents: [ + { + name: "helper", + description: "Helps with work", + systemPrompt: "Help.", + }, + ], + }); + + const middleware = getMiddlewareStack("helper"); + expect(middleware).not.toContain(custom); + }); + + it("replaces declarative subagent defaults with same-name spec middleware", () => { + const custom = createCustomSummarizationMiddleware(); + + createDeepAgent({ + model: fakeModel, + name: "main", + subagents: [ + { + name: "helper", + description: "Helps with work", + systemPrompt: "Help.", + middleware: [custom], + }, + ], + }); + + const middleware = getMiddlewareStack("helper"); + const summarization = middleware.filter( + (entry) => entry.name === "SummarizationMiddleware", + ); + expect(summarization).toHaveLength(1); + expect(summarization[0]).toBe(custom); + }); +}); diff --git a/libs/deepagents/src/middleware/tool_exclusion.ts b/libs/deepagents/src/middleware/tool_exclusion.ts new file mode 100644 index 000000000..efeff2281 --- /dev/null +++ b/libs/deepagents/src/middleware/tool_exclusion.ts @@ -0,0 +1,32 @@ +import { createMiddleware, type AgentMiddleware } from "langchain"; + +function hasToolName(tool: unknown): tool is { name: string } { + return ( + tool !== null && + typeof tool === "object" && + "name" in tool && + typeof tool.name === "string" + ); +} + +/** + * Create middleware that removes excluded tools after all tool-injecting + * middleware has had a chance to add tools to the request. + * + * @internal + */ +export function createToolExclusionMiddleware( + excludedTools: ReadonlySet, +): AgentMiddleware { + return createMiddleware({ + name: "_ToolExclusionMiddleware", + wrapModelCall(request, handler) { + return handler({ + ...request, + tools: request.tools?.filter( + (tool) => !hasToolName(tool) || !excludedTools.has(tool.name), + ), + }); + }, + }); +} diff --git a/libs/deepagents/src/middleware/utils.ts b/libs/deepagents/src/middleware/utils.ts index 1612954d9..751264c2a 100644 --- a/libs/deepagents/src/middleware/utils.ts +++ b/libs/deepagents/src/middleware/utils.ts @@ -5,6 +5,76 @@ */ import { SystemMessage } from "@langchain/core/messages"; +import type { AgentMiddleware } from "langchain"; + +/** + * Merge custom middleware into an assembled stack by `.name`. + * + * Matching custom middleware replaces the existing entry in place. New + * middleware is appended after the base stack in caller-provided order. + */ +export function mergeMiddleware( + base: readonly AgentMiddleware[], + custom: readonly AgentMiddleware[], +): AgentMiddleware[] { + const merged = new Map( + base.map((middleware) => [middleware.name, middleware]), + ); + for (const middleware of custom) { + merged.set(middleware.name, middleware); + } + return [...merged.values()]; +} + +function middlewareNames(middleware: readonly AgentMiddleware[]): Set { + return new Set(middleware.map((entry) => entry.name)); +} + +function matchingMiddleware( + middleware: readonly AgentMiddleware[], + names: ReadonlySet, +): AgentMiddleware[] { + return middleware.filter((entry) => names.has(entry.name)); +} + +/** + * Merge custom middleware into default and tail middleware segments. + * + * Same-name custom entries replace matching defaults in either segment. Novel + * custom entries are inserted between the default and tail segments unless + * `appendNew` is false. + */ +export function mergeMiddlewareStack( + defaultMiddleware: readonly AgentMiddleware[], + customMiddleware: readonly AgentMiddleware[], + tailMiddleware: readonly AgentMiddleware[] = [], + options: { appendNew?: boolean } = {}, +): AgentMiddleware[] { + const defaultMiddlewareNames = middlewareNames(defaultMiddleware); + const tailMiddlewareNames = middlewareNames(tailMiddleware); + const knownMiddlewareNames = new Set([ + ...defaultMiddlewareNames, + ...tailMiddlewareNames, + ]); + const novelMiddleware = + options.appendNew === false + ? [] + : customMiddleware.filter( + (entry) => !knownMiddlewareNames.has(entry.name), + ); + + return [ + ...mergeMiddleware( + defaultMiddleware, + matchingMiddleware(customMiddleware, defaultMiddlewareNames), + ), + ...novelMiddleware, + ...mergeMiddleware( + tailMiddleware, + matchingMiddleware(customMiddleware, tailMiddlewareNames), + ), + ]; +} /** * Append text to a system message.