mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
feat: memory injection (#2631)
This commit is contained in:
@@ -89,6 +89,10 @@
|
||||
"types": "./dist/src/vended-plugins/context-offloader/index.d.ts",
|
||||
"default": "./dist/src/vended-plugins/context-offloader/index.js"
|
||||
},
|
||||
"./vended-plugins/context-injector": {
|
||||
"types": "./dist/src/vended-plugins/context-injector/index.d.ts",
|
||||
"default": "./dist/src/vended-plugins/context-injector/index.js"
|
||||
},
|
||||
"./vended-plugins/goal": {
|
||||
"types": "./dist/src/vended-plugins/goal/index.d.ts",
|
||||
"default": "./dist/src/vended-plugins/goal/index.js"
|
||||
|
||||
@@ -346,6 +346,10 @@ export type {
|
||||
MemoryToolConfig,
|
||||
MemoryAddToolConfig,
|
||||
MemoryManagerConfig,
|
||||
MemoryInjectionConfig,
|
||||
InjectionConfig,
|
||||
InjectionTrigger,
|
||||
InjectionContext,
|
||||
} from './memory/index.js'
|
||||
export { ExtractionTrigger, InvocationTrigger, IntervalTrigger, ModelExtractor } from './memory/index.js'
|
||||
export type {
|
||||
|
||||
@@ -0,0 +1,240 @@
|
||||
import { describe, it, expect, vi } from 'vitest'
|
||||
import { Message, TextBlock, ToolResultBlock } from '../../types/messages.js'
|
||||
import type { MessageData } from '../../types/messages.js'
|
||||
import { foldIntoLastUserMessage, isUserTurn, resolveTrigger, createInjectionMiddleware } from '../message-injection.js'
|
||||
import type { InvokeModelContext } from '../../middleware/index.js'
|
||||
import type { InjectionContext } from '../types.js'
|
||||
import { createMockAgent } from '../../__fixtures__/agent-helpers.js'
|
||||
import { logger } from '../../logging/logger.js'
|
||||
|
||||
const user = (text: string) => new Message({ role: 'user', content: [new TextBlock(text)] })
|
||||
const assistant = (text: string) => new Message({ role: 'assistant', content: [new TextBlock(text)] })
|
||||
const toolResult = () =>
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new ToolResultBlock({ toolUseId: 't1', status: 'success', content: [new TextBlock('done')] })],
|
||||
})
|
||||
|
||||
// resolveTrigger predicates take an InjectionContext; tests only exercise `messages`, so a minimal bag suffices.
|
||||
const injectionCtx = (messages: MessageData[]) => ({ messages }) as unknown as InjectionContext
|
||||
describe('foldIntoLastUserMessage', () => {
|
||||
it('prepends the text as a leading TextBlock on the last user message, ahead of its content', () => {
|
||||
const messages = [user('original task'), assistant('prior step'), user('next ask')]
|
||||
const result = foldIntoLastUserMessage(messages, 'INJECTED')
|
||||
|
||||
// The earlier user/assistant turns are untouched; the last user message gains a leading INJECTED
|
||||
// block ahead of its own content, keeping the user's ask in the recency slot.
|
||||
expect(result.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'user', content: [{ text: 'original task' }] },
|
||||
{ role: 'assistant', content: [{ text: 'prior step' }] },
|
||||
{ role: 'user', content: [{ text: 'INJECTED' }, { text: 'next ask' }] },
|
||||
])
|
||||
})
|
||||
|
||||
it('returns a new array and does not mutate the input or its messages', () => {
|
||||
const original = user('ask')
|
||||
const messages = [assistant('prior'), original]
|
||||
const result = foldIntoLastUserMessage(messages, 'INJECTED')
|
||||
|
||||
expect(result).not.toBe(messages)
|
||||
expect(messages[1]).toBe(original)
|
||||
expect(original.content).toHaveLength(1) // untouched
|
||||
expect(result[1]).not.toBe(original)
|
||||
})
|
||||
|
||||
it('appends after the tool result block when the target is a tool-result turn', () => {
|
||||
const tr = toolResult()
|
||||
const result = foldIntoLastUserMessage([user('task'), assistant('thinking'), tr], 'INJECTED')
|
||||
|
||||
// Providers require the tool result to be the first block in the turn, so the injected text is
|
||||
// appended rather than prepended here.
|
||||
expect(result.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'user', content: [{ text: 'task' }] },
|
||||
{ role: 'assistant', content: [{ text: 'thinking' }] },
|
||||
{ role: 'user', content: [tr.toJSON().content[0], { text: 'INJECTED' }] },
|
||||
])
|
||||
})
|
||||
|
||||
it('targets the most recent user message when several exist', () => {
|
||||
const messages = [user('first'), assistant('a'), user('second')]
|
||||
const result = foldIntoLastUserMessage(messages, 'INJECTED')
|
||||
|
||||
expect(result.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'user', content: [{ text: 'first' }] }, // earlier user untouched
|
||||
{ role: 'assistant', content: [{ text: 'a' }] },
|
||||
{ role: 'user', content: [{ text: 'INJECTED' }, { text: 'second' }] },
|
||||
])
|
||||
})
|
||||
|
||||
it('preserves message metadata on the folded message', () => {
|
||||
const tagged = new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('ask')],
|
||||
metadata: { custom: { keep: 'me' } },
|
||||
})
|
||||
const result = foldIntoLastUserMessage([tagged], 'INJECTED')
|
||||
expect(result[0]!.metadata?.custom).toStrictEqual({ keep: 'me' })
|
||||
})
|
||||
|
||||
it('returns the input unchanged when there is no user message', () => {
|
||||
const messages = [assistant('only assistant')]
|
||||
const result = foldIntoLastUserMessage(messages, 'INJECTED')
|
||||
expect(result).toBe(messages)
|
||||
})
|
||||
})
|
||||
|
||||
describe('isUserTurn', () => {
|
||||
it('is true when the last message is a plain user ask', () => {
|
||||
expect(isUserTurn([assistant('prior').toJSON(), user('ask').toJSON()])).toBe(true)
|
||||
})
|
||||
|
||||
it('is false when the last message is a user tool-result turn', () => {
|
||||
expect(isUserTurn([user('task').toJSON(), assistant('a').toJSON(), toolResult().toJSON()])).toBe(false)
|
||||
})
|
||||
|
||||
it('is false when the last message is an assistant message', () => {
|
||||
expect(isUserTurn([user('ask').toJSON(), assistant('reply').toJSON()])).toBe(false)
|
||||
})
|
||||
|
||||
it('is false for an empty conversation', () => {
|
||||
expect(isUserTurn([])).toBe(false)
|
||||
})
|
||||
})
|
||||
|
||||
describe('resolveTrigger', () => {
|
||||
it('defaults (undefined) to the userTurn policy', () => {
|
||||
const trigger = resolveTrigger(undefined)
|
||||
expect(trigger(injectionCtx([user('ask').toJSON()]))).toBe(true)
|
||||
expect(trigger(injectionCtx([toolResult().toJSON()]))).toBe(false)
|
||||
})
|
||||
|
||||
it("'userTurn' uses isUserTurn", () => {
|
||||
const trigger = resolveTrigger('userTurn')
|
||||
expect(trigger(injectionCtx([user('ask').toJSON()]))).toBe(true)
|
||||
expect(trigger(injectionCtx([toolResult().toJSON()]))).toBe(false)
|
||||
})
|
||||
|
||||
it("'everyTurn' always fires", () => {
|
||||
const trigger = resolveTrigger('everyTurn')
|
||||
expect(trigger(injectionCtx([]))).toBe(true)
|
||||
expect(trigger(injectionCtx([toolResult().toJSON()]))).toBe(true)
|
||||
})
|
||||
|
||||
it('uses a custom predicate over the context', () => {
|
||||
const trigger = resolveTrigger((context) => context.messages.length >= 2)
|
||||
expect(trigger(injectionCtx([user('a').toJSON()]))).toBe(false)
|
||||
expect(trigger(injectionCtx([user('a').toJSON(), assistant('b').toJSON()]))).toBe(true)
|
||||
})
|
||||
|
||||
it('fails open (returns false, logs) when a custom predicate throws', () => {
|
||||
const warn = vi.spyOn(logger, 'warn').mockImplementation(() => {})
|
||||
const trigger = resolveTrigger(() => {
|
||||
throw new Error('boom')
|
||||
})
|
||||
expect(trigger(injectionCtx([user('ask').toJSON()]))).toBe(false)
|
||||
expect(warn).toHaveBeenCalled()
|
||||
warn.mockRestore()
|
||||
})
|
||||
})
|
||||
|
||||
describe('createInjectionMiddleware', () => {
|
||||
// The handler is an InvokeModelStage.Input transformer. It reads `context.messages` and derives the
|
||||
// InjectionContext (appState/agent) from `context.agent`, then spreads the rest through, so a context
|
||||
// carrying `messages` plus a mock agent exercises it faithfully.
|
||||
const ctx = (messages: Message[]) => ({ messages, agent: createMockAgent() }) as unknown as InvokeModelContext
|
||||
|
||||
it('folds renderContent() text into the latest user message, leaving other context fields intact', async () => {
|
||||
const handler = createInjectionMiddleware({ renderContent: async () => 'INJECTED' })
|
||||
const result = await handler(ctx([assistant('prior'), user('ask')]))
|
||||
|
||||
expect(result.messages.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'assistant', content: [{ text: 'prior' }] },
|
||||
{ role: 'user', content: [{ text: 'INJECTED' }, { text: 'ask' }] },
|
||||
])
|
||||
})
|
||||
|
||||
it('passes an InjectionContext carrying the conversation (as data) to renderContent', async () => {
|
||||
const seen: string[] = []
|
||||
const handler = createInjectionMiddleware({
|
||||
renderContent: async (context) => {
|
||||
seen.push(...context.messages.map((m) => m.role))
|
||||
return 'x'
|
||||
},
|
||||
})
|
||||
await handler(ctx([assistant('prior'), user('ask')]))
|
||||
|
||||
expect(seen).toStrictEqual(['assistant', 'user'])
|
||||
})
|
||||
|
||||
it('exposes appState and the agent on the InjectionContext', async () => {
|
||||
const appState = { get: () => 'stashed' }
|
||||
const agent = { appState } as unknown as InvokeModelContext['agent']
|
||||
const input = { messages: [user('ask')], agent } as unknown as InvokeModelContext
|
||||
let received: { appState: unknown; agent: unknown } | undefined
|
||||
const handler = createInjectionMiddleware({
|
||||
renderContent: async (context) => {
|
||||
received = { appState: context.appState, agent: context.agent }
|
||||
return undefined
|
||||
},
|
||||
})
|
||||
await handler(input)
|
||||
|
||||
expect(received).toStrictEqual({ appState, agent })
|
||||
})
|
||||
|
||||
it('returns the context unchanged when the trigger does not fire', async () => {
|
||||
const renderContent = vi.fn(async () => 'x')
|
||||
const handler = createInjectionMiddleware({ renderContent }) // default 'userTurn'
|
||||
const input = ctx([user('task'), assistant('a'), toolResult()])
|
||||
const result = await handler(input)
|
||||
|
||||
expect(result).toBe(input)
|
||||
expect(renderContent).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("'everyTurn' injects on an autonomous tool-result turn, keeping the tool result first", async () => {
|
||||
const handler = createInjectionMiddleware({ trigger: 'everyTurn', renderContent: async () => 'INJECTED' })
|
||||
const tr = toolResult()
|
||||
const result = await handler(ctx([user('task'), assistant('a'), tr]))
|
||||
|
||||
// The most recent user message on a tool-result turn carries the tool result, which must stay the
|
||||
// first block, so the injected text is appended after it.
|
||||
expect(result.messages.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'user', content: [{ text: 'task' }] },
|
||||
{ role: 'assistant', content: [{ text: 'a' }] },
|
||||
{ role: 'user', content: [tr.toJSON().content[0], { text: 'INJECTED' }] },
|
||||
])
|
||||
})
|
||||
|
||||
it('returns the context unchanged when renderContent yields empty text', async () => {
|
||||
const handler = createInjectionMiddleware({ renderContent: async () => ' ' })
|
||||
const input = ctx([assistant('prior'), user('ask')])
|
||||
const result = await handler(input)
|
||||
|
||||
expect(result).toBe(input)
|
||||
})
|
||||
|
||||
it('fails open (returns the context unchanged, logs) when renderContent throws', async () => {
|
||||
const warn = vi.spyOn(logger, 'warn').mockImplementation(() => {})
|
||||
const handler = createInjectionMiddleware({
|
||||
renderContent: async () => {
|
||||
throw new Error('boom')
|
||||
},
|
||||
})
|
||||
const input = ctx([assistant('prior'), user('ask')])
|
||||
const result = await handler(input)
|
||||
|
||||
expect(result).toBe(input)
|
||||
expect(warn).toHaveBeenCalled()
|
||||
warn.mockRestore()
|
||||
})
|
||||
|
||||
it('does not mutate the original context messages', async () => {
|
||||
const handler = createInjectionMiddleware({ renderContent: async () => 'INJECTED' })
|
||||
const input = ctx([assistant('prior'), user('ask')])
|
||||
const before = input.messages[1]!
|
||||
await handler(input)
|
||||
|
||||
expect(before.content).toHaveLength(1) // the original user message is untouched
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,5 @@
|
||||
/**
|
||||
* Configuration types for context injection, shared by the `ContextInjector` plugin and
|
||||
* `MemoryManager`'s `injection` config.
|
||||
*/
|
||||
export type { InjectionConfig, InjectionTrigger, InjectionContext } from './types.js'
|
||||
@@ -0,0 +1,146 @@
|
||||
import { Message, TextBlock } from '../types/messages.js'
|
||||
import type { MessageData } from '../types/messages.js'
|
||||
import { logger } from '../logging/logger.js'
|
||||
import { normalizeError } from '../errors.js'
|
||||
import type { InvokeModelContext } from '../middleware/index.js'
|
||||
import type { InjectionMiddlewareOptions, InjectionTrigger, InjectionContext } from './types.js'
|
||||
|
||||
/**
|
||||
* Builds an `InvokeModelStage` `Input` handler that folds {@link InjectionMiddlewareOptions.renderContent}'s
|
||||
* text into the latest user message, ephemerally — the model sees the augmented input for this one call
|
||||
* while the agent's durable history is never touched.
|
||||
*
|
||||
* Runs as an input-phase transformer (`(ctx) => ctx`): it gates on the resolved trigger, asks
|
||||
* `renderContent` for the text, and returns a context with the folded messages. Anything that skips —
|
||||
* the trigger not firing, `renderContent` returning empty, or any callback throwing — returns the
|
||||
* context unchanged so the model call proceeds (fail open). The injected text never enters durable
|
||||
* history because the input phase only rewrites the per-call context, not the agent's stored messages.
|
||||
*
|
||||
* @param opts - The trigger and `renderContent` callback the handler uses
|
||||
* @returns An `InvokeModelStage.Input` handler that returns a (possibly) folded context
|
||||
* @internal Delivery primitive. Reach injection through `ContextInjector` or `MemoryManager`.
|
||||
*/
|
||||
export function createInjectionMiddleware(
|
||||
opts: InjectionMiddlewareOptions
|
||||
): (context: InvokeModelContext) => Promise<InvokeModelContext> {
|
||||
const trigger = resolveTrigger(opts.trigger)
|
||||
return async (context) => {
|
||||
const agent = context.agent
|
||||
const injectionContext: InjectionContext = {
|
||||
messages: context.messages.map((message) => message.toJSON()),
|
||||
appState: agent.appState,
|
||||
agent,
|
||||
}
|
||||
if (!trigger(injectionContext)) {
|
||||
return context
|
||||
}
|
||||
|
||||
let text: string | undefined
|
||||
try {
|
||||
text = await opts.renderContent(injectionContext)
|
||||
} catch (error) {
|
||||
logger.warn(`reason=<${normalizeError(error).message}> | injection renderContent threw; skipping injection`)
|
||||
return context
|
||||
}
|
||||
if (!text?.trim()) {
|
||||
return context
|
||||
}
|
||||
|
||||
return { ...context, messages: foldIntoLastUserMessage([...context.messages], text) }
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Resolves an {@link InjectionTrigger} name or predicate into a single gate predicate over the
|
||||
* {@link InjectionContext}.
|
||||
*
|
||||
* `'userTurn'` maps to {@link isUserTurn} (over `ctx.messages`); `'everyTurn'` to an always-true gate;
|
||||
* a user-supplied predicate is wrapped so that a throw fails open (logs and skips injection rather than
|
||||
* aborting the model call).
|
||||
*
|
||||
* @param trigger - An {@link InjectionTrigger} name, a predicate, or `undefined` (defaults to `'userTurn'`)
|
||||
* @returns A predicate that, given the {@link InjectionContext}, returns whether to inject this call
|
||||
* @internal Delivery primitive. Reach injection through `ContextInjector` or `MemoryManager`.
|
||||
*/
|
||||
export function resolveTrigger(
|
||||
trigger: InjectionTrigger | ((context: InjectionContext) => boolean) | undefined
|
||||
): (context: InjectionContext) => boolean {
|
||||
if (trigger === undefined || trigger === 'userTurn') {
|
||||
return (context) => isUserTurn(context.messages)
|
||||
}
|
||||
if (trigger === 'everyTurn') {
|
||||
return () => true
|
||||
}
|
||||
const predicate = trigger
|
||||
return (context) => {
|
||||
try {
|
||||
return predicate(context)
|
||||
} catch (error) {
|
||||
logger.warn(`reason=<${normalizeError(error).message}> | injection trigger threw; skipping injection`)
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Whether the latest message is a fresh user ask: a `user` message carrying no tool result. This is
|
||||
* the `'userTurn'` policy — it distinguishes a new chat ask from an autonomous tool-result turn.
|
||||
*
|
||||
* @param messages - The current conversation, as data
|
||||
* @returns `true` when the latest message is a plain user ask, otherwise `false`
|
||||
* @internal Delivery primitive. Reach injection through `ContextInjector` or `MemoryManager`.
|
||||
*/
|
||||
export function isUserTurn(messages: MessageData[]): boolean {
|
||||
const last = messages[messages.length - 1]
|
||||
return !!last && last.role === 'user' && !last.content.some((block) => 'toolResult' in block)
|
||||
}
|
||||
|
||||
/**
|
||||
* Folds `text` into the most recent `user` message as a {@link TextBlock}, returning a NEW array. Other
|
||||
* messages are returned as-is.
|
||||
*
|
||||
* Folding into the existing user message (rather than inserting a standalone message) keeps role
|
||||
* alternation valid in both chat and the autonomous tool loop. The block is placed to keep the message
|
||||
* valid for the model:
|
||||
* - A plain user ask: the text is **prepended**, leaving the user's own ask in the recency slot — the
|
||||
* last thing the model reads.
|
||||
* - A tool-result turn (the message carries a `ToolResultBlock`): the text is **appended**,
|
||||
* because providers require the tool result to be the first content block in the turn that answers a
|
||||
* tool use.
|
||||
*
|
||||
* {@link Message} fields are readonly, so the target is rebuilt as a new {@link Message}. When there is
|
||||
* no `user` message, the input array is returned unchanged.
|
||||
*
|
||||
* @param messages - The conversation to fold into
|
||||
* @param text - The text to fold into the most recent user message
|
||||
* @returns A new array with the folded message, or the input array when there is no user message
|
||||
* @internal Delivery primitive. Reach injection through `ContextInjector` or `MemoryManager`.
|
||||
*/
|
||||
export function foldIntoLastUserMessage(messages: Message[], text: string): Message[] {
|
||||
let targetIndex = -1
|
||||
for (let i = messages.length - 1; i >= 0; i--) {
|
||||
if (messages[i]!.role === 'user') {
|
||||
targetIndex = i
|
||||
break
|
||||
}
|
||||
}
|
||||
if (targetIndex < 0) {
|
||||
return messages
|
||||
}
|
||||
|
||||
const target = messages[targetIndex]!
|
||||
const injected = new TextBlock(text)
|
||||
// A tool result must stay the first block in the turn that answers a tool use, so append rather than
|
||||
// prepend when the target carries one.
|
||||
const hasToolResult = target.content.some((block) => block.type === 'toolResultBlock')
|
||||
const content = hasToolResult ? [...target.content, injected] : [injected, ...target.content]
|
||||
const folded = new Message({
|
||||
role: target.role,
|
||||
content,
|
||||
...(target.metadata !== undefined && { metadata: target.metadata }),
|
||||
})
|
||||
|
||||
const result = [...messages]
|
||||
result[targetIndex] = folded
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
import type { MessageData } from '../types/messages.js'
|
||||
import type { LocalAgent } from '../types/agent.js'
|
||||
import type { StateStore } from '../state-store.js'
|
||||
|
||||
/**
|
||||
* Determines when injection runs before a model call.
|
||||
*
|
||||
* - `'userTurn'`: only when the latest message is a fresh user ask (a `user` message with no tool
|
||||
* result) — the common case for chat agents, where it keeps the user's ask the final message the
|
||||
* model sees.
|
||||
* - `'everyTurn'`: before every model call, including mid-task tool-result turns — for autonomous
|
||||
* agents that should consult injected context at each step.
|
||||
*
|
||||
* For finer control, pass a predicate as {@link InjectionConfig.trigger} instead.
|
||||
*/
|
||||
export type InjectionTrigger = 'userTurn' | 'everyTurn'
|
||||
|
||||
/**
|
||||
* The context an injection consumer receives on each model call, passed to `renderContent` and to a
|
||||
* predicate {@link InjectionConfig.trigger}.
|
||||
*/
|
||||
export interface InjectionContext {
|
||||
/** The current conversation, as data. */
|
||||
messages: MessageData[]
|
||||
/** Durable app state shared across calls, hooks, and tools — read what a tool stashed last turn. */
|
||||
appState: StateStore
|
||||
/** The agent the injection is attached to (escape hatch for advanced consumers). */
|
||||
agent: LocalAgent
|
||||
}
|
||||
|
||||
/**
|
||||
* Configuration common to every injection consumer: when to inject. What text to inject is a consumer
|
||||
* concern, added by the interfaces that extend this one (e.g. {@link MemoryInjectionConfig}).
|
||||
*/
|
||||
export interface InjectionConfig {
|
||||
/**
|
||||
* When injection runs. An {@link InjectionTrigger} name selects a built-in policy; a predicate is
|
||||
* the escape hatch — it receives the {@link InjectionContext} and returns whether to inject this
|
||||
* call. A predicate that throws fails open (injection is skipped, the model call proceeds).
|
||||
*
|
||||
* @defaultValue 'userTurn'
|
||||
*/
|
||||
trigger?: InjectionTrigger | ((context: InjectionContext) => boolean)
|
||||
}
|
||||
|
||||
/**
|
||||
* Options for {@link createInjectionMiddleware}.
|
||||
*
|
||||
* The engine is text-in: it knows nothing about queries, search, or rendering. A consumer supplies a
|
||||
* single {@link InjectionMiddlewareOptions.renderContent} callback that returns the text to fold into
|
||||
* the conversation, and (optionally) a trigger that gates when to do so.
|
||||
*
|
||||
* @internal Engine options. Consumers configure injection via `ContextInjectorConfig` or
|
||||
* `MemoryInjectionConfig`, not this type.
|
||||
*/
|
||||
export interface InjectionMiddlewareOptions {
|
||||
/**
|
||||
* When to inject. See {@link InjectionConfig.trigger}. Defaults to `'userTurn'`.
|
||||
*/
|
||||
trigger?: InjectionTrigger | ((context: InjectionContext) => boolean)
|
||||
/**
|
||||
* Returns the text to fold into the latest user message, or `undefined`/`''` to skip this call. A
|
||||
* callback that throws fails open (injection is skipped, the model call proceeds).
|
||||
*/
|
||||
renderContent: (context: InjectionContext) => Promise<string | undefined>
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
/**
|
||||
* Minimal XML escaping for folding untrusted text into an XML-shaped block.
|
||||
*
|
||||
* Memory entries and other injected content are frequently user-derived, so interpolating them raw
|
||||
* into `<entry>…</entry>` both breaks the block structurally (a stray `</entry>` or `"`) and opens a
|
||||
* stored-prompt-injection surface. These helpers are deliberately tiny — enough to keep a `<memory>`
|
||||
* block well-formed, not a general-purpose serializer.
|
||||
*/
|
||||
|
||||
/**
|
||||
* Escapes text content for placement between XML tags: `&` (first, so later replacements are not
|
||||
* double-escaped), then `<` and `>`.
|
||||
*
|
||||
* @param value - The raw text to escape
|
||||
* @returns The escaped text, safe to place in element content
|
||||
* @internal Used by the memory default formatter
|
||||
*/
|
||||
export function escapeXmlText(value: string): string {
|
||||
return value.replace(/&/g, '&').replace(/</g, '<').replace(/>/g, '>')
|
||||
}
|
||||
|
||||
/**
|
||||
* Escapes a value for placement inside a double-quoted XML attribute: the {@link escapeXmlText} rules
|
||||
* plus `"` and `'`.
|
||||
*
|
||||
* @param value - The raw attribute value to escape
|
||||
* @returns The escaped value, safe to place inside a quoted attribute
|
||||
* @internal Default-format helper; not part of the public surface.
|
||||
*/
|
||||
export function escapeXmlAttr(value: string): string {
|
||||
return escapeXmlText(value).replace(/"/g, '"').replace(/'/g, ''')
|
||||
}
|
||||
@@ -6,6 +6,11 @@ import { tool } from '../../tools/tool-factory.js'
|
||||
import type { MemoryStore, MemoryEntry } from '../types.js'
|
||||
import type { InvokableTool, Tool } from '../../tools/tool.js'
|
||||
import { logger } from '../../logging/logger.js'
|
||||
import { Message, TextBlock, ToolUseBlock, ToolResultBlock } from '../../types/messages.js'
|
||||
import type { MessageData } from '../../types/messages.js'
|
||||
import { InvokeModelStage } from '../../middleware/index.js'
|
||||
import type { InvokeModelContext } from '../../middleware/index.js'
|
||||
import { createMockAgent } from '../../__fixtures__/agent-helpers.js'
|
||||
|
||||
function createMockStore(
|
||||
name: string,
|
||||
@@ -611,6 +616,239 @@ describe('MemoryManager', () => {
|
||||
const mm = new MemoryManager({ stores: [createMockStore('test')] })
|
||||
expect(() => mm.initAgent({} as any)).not.toThrow()
|
||||
})
|
||||
|
||||
it('does not register injection middleware when injection is disabled', () => {
|
||||
const mm = new MemoryManager({ stores: [createMockStore('test')] })
|
||||
const addMiddleware = vi.fn()
|
||||
mm.initAgent(createMockAgent({ extra: { addMiddleware } as never }))
|
||||
expect(addMiddleware).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('registers an InvokeModelStage input middleware when injection is enabled', () => {
|
||||
const mm = new MemoryManager({ stores: [createMockStore('test')], injection: true })
|
||||
const addMiddleware = vi.fn()
|
||||
mm.initAgent(createMockAgent({ extra: { addMiddleware } as never }))
|
||||
|
||||
expect(addMiddleware).toHaveBeenCalledTimes(1)
|
||||
expect(addMiddleware.mock.calls[0]![0]).toBe(InvokeModelStage.Input)
|
||||
expect(typeof addMiddleware.mock.calls[0]![1]).toBe('function')
|
||||
})
|
||||
|
||||
it('wires the registered middleware to the memory provide pipeline (folds a search hit)', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'dark mode preferred' }] })
|
||||
const mm = new MemoryManager({ stores: [store], injection: true })
|
||||
const addMiddleware = vi.fn()
|
||||
const agent = createMockAgent({ extra: { addMiddleware } as never })
|
||||
mm.initAgent(agent)
|
||||
|
||||
const handler = addMiddleware.mock.calls[0]![1] as (ctx: InvokeModelContext) => Promise<InvokeModelContext>
|
||||
const messages = [
|
||||
new Message({ role: 'assistant', content: [new TextBlock('prior')] }),
|
||||
new Message({ role: 'user', content: [new TextBlock('what is my plan')] }),
|
||||
]
|
||||
const result = await handler({ messages, agent } as unknown as InvokeModelContext)
|
||||
|
||||
expect(result.messages.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'assistant', content: [{ text: 'prior' }] },
|
||||
{
|
||||
role: 'user',
|
||||
content: [
|
||||
{ text: '<memory>\n<entry source="s">dark mode preferred</entry>\n</memory>' },
|
||||
{ text: 'what is my plan' },
|
||||
],
|
||||
},
|
||||
])
|
||||
expect(store.search).toHaveBeenCalledWith('what is my plan', { maxSearchResults: 5 })
|
||||
})
|
||||
})
|
||||
|
||||
// The injection delivery (folding text into the model input) is wired through the InvokeModelStage
|
||||
// input middleware (see the `initAgent` tests). The memory-owned `provide` pipeline below — query
|
||||
// derivation, search, and formatting — is exercised directly via `_provideMemoryContext`, the
|
||||
// callback the middleware invokes.
|
||||
describe('injection', () => {
|
||||
const assistant = (text: string) => new Message({ role: 'assistant', content: [new TextBlock(text)] })
|
||||
const user = (text: string) => new Message({ role: 'user', content: [new TextBlock(text)] })
|
||||
const toolUse = () =>
|
||||
new Message({ role: 'assistant', content: [new ToolUseBlock({ name: 'x', toolUseId: 't1', input: {} })] })
|
||||
const toolResult = () =>
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new ToolResultBlock({ toolUseId: 't1', status: 'success', content: [new TextBlock('done')] })],
|
||||
})
|
||||
|
||||
// Calls the (private) provide pipeline with the manager's resolved injection config.
|
||||
function provide(mm: MemoryManager, messages: Message[]): Promise<string | undefined> {
|
||||
const data = messages.map((m) => m.toJSON())
|
||||
const config = (mm as unknown as { _injectionConfig: object | false })._injectionConfig
|
||||
return (
|
||||
mm as unknown as { _provideMemoryContext(m: MessageData[], c: object): Promise<string | undefined> }
|
||||
)._provideMemoryContext(data, config === false ? {} : config)
|
||||
}
|
||||
|
||||
describe('config resolution', () => {
|
||||
const injectionConfig = (mm: MemoryManager) =>
|
||||
(mm as unknown as { _injectionConfig: object | false })._injectionConfig
|
||||
|
||||
it('defaults to false (disabled) when injection is omitted', () => {
|
||||
expect(injectionConfig(new MemoryManager({ stores: [createMockStore('s')] }))).toBe(false)
|
||||
})
|
||||
|
||||
it('is false when injection is explicitly false', () => {
|
||||
expect(injectionConfig(new MemoryManager({ stores: [createMockStore('s')], injection: false }))).toBe(false)
|
||||
})
|
||||
|
||||
it('resolves to an empty config when injection is true', () => {
|
||||
expect(injectionConfig(new MemoryManager({ stores: [createMockStore('s')], injection: true }))).toStrictEqual(
|
||||
{}
|
||||
)
|
||||
})
|
||||
|
||||
it('passes an injection config object through unchanged', () => {
|
||||
const cfg = { maxEntries: 5 }
|
||||
expect(injectionConfig(new MemoryManager({ stores: [createMockStore('s')], injection: cfg }))).toBe(cfg)
|
||||
})
|
||||
})
|
||||
|
||||
describe('query', () => {
|
||||
it('uses the latest user ask on a user turn (adaptive default)', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'fact' }] })
|
||||
const mm = new MemoryManager({ stores: [store], injection: true })
|
||||
|
||||
await provide(mm, [assistant('prior step'), user('what is my plan')])
|
||||
|
||||
expect(store.search).toHaveBeenCalledWith('what is my plan', { maxSearchResults: 5 })
|
||||
})
|
||||
|
||||
it('uses the most recent assistant text on an autonomous (tool-result) turn', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'fact' }] })
|
||||
const mm = new MemoryManager({ stores: [store], injection: true })
|
||||
|
||||
await provide(mm, [user('task'), assistant('the previous step result'), toolResult()])
|
||||
|
||||
expect(store.search).toHaveBeenCalledWith('the previous step result', { maxSearchResults: 5 })
|
||||
})
|
||||
|
||||
it('honors a custom query', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'fact' }] })
|
||||
const mm = new MemoryManager({ stores: [store], injection: { query: () => 'custom query' } })
|
||||
|
||||
await provide(mm, [assistant('prior'), user('ask')])
|
||||
|
||||
expect(store.search).toHaveBeenCalledWith('custom query', { maxSearchResults: 5 })
|
||||
})
|
||||
|
||||
it('skips (returns undefined) when a custom query returns undefined', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'fact' }] })
|
||||
const mm = new MemoryManager({ stores: [store], injection: { query: () => undefined } })
|
||||
|
||||
await expect(provide(mm, [assistant('prior'), user('ask')])).resolves.toBeUndefined()
|
||||
expect(store.search).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('fails open (returns undefined) when a custom query throws', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'fact' }] })
|
||||
const mm = new MemoryManager({
|
||||
stores: [store],
|
||||
injection: {
|
||||
query: () => {
|
||||
throw new Error('boom')
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
await expect(provide(mm, [assistant('prior'), user('ask')])).resolves.toBeUndefined()
|
||||
expect(store.search).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('skips when the latest assistant message has no text (pure tool use)', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'fact' }] })
|
||||
const mm = new MemoryManager({ stores: [store], injection: true })
|
||||
|
||||
await expect(provide(mm, [toolUse(), toolResult()])).resolves.toBeUndefined()
|
||||
expect(store.search).not.toHaveBeenCalled()
|
||||
})
|
||||
})
|
||||
|
||||
describe('search', () => {
|
||||
it('returns undefined when search yields no entries', async () => {
|
||||
const store = createMockStore('s', { entries: [] })
|
||||
const mm = new MemoryManager({ stores: [store], injection: true })
|
||||
|
||||
await expect(provide(mm, [assistant('prior'), user('ask')])).resolves.toBeUndefined()
|
||||
expect(store.search).toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('honors maxEntries and caps the rendered entries', async () => {
|
||||
const store = createMockStore('s', {
|
||||
entries: [{ content: 'A' }, { content: 'B' }, { content: 'C' }],
|
||||
})
|
||||
const mm = new MemoryManager({
|
||||
stores: [store],
|
||||
injection: { maxEntries: 2, format: ({ entries }) => entries.map((e) => e.content).join(',') },
|
||||
})
|
||||
|
||||
const text = await provide(mm, [assistant('prior'), user('ask')])
|
||||
|
||||
expect(store.search).toHaveBeenCalledWith('ask', { maxSearchResults: 2 })
|
||||
expect(text).toBe('A,B')
|
||||
})
|
||||
})
|
||||
|
||||
describe('format', () => {
|
||||
it('default renders a <memory> block with per-entry source attribution', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'dark mode preferred' }] })
|
||||
const mm = new MemoryManager({ stores: [store], injection: true })
|
||||
|
||||
// search() stamps storeName onto each entry, so the default format attributes the source.
|
||||
const text = await provide(mm, [assistant('prior'), user('ask')])
|
||||
expect(text).toBe('<memory>\n<entry source="s">dark mode preferred</entry>\n</memory>')
|
||||
})
|
||||
|
||||
it('omits the source attribute for an entry with no storeName', () => {
|
||||
const mm = new MemoryManager({ stores: [createMockStore('s')], injection: true })
|
||||
const text = (mm as unknown as { _defaultInjectionFormat(e: MemoryEntry[]): string })._defaultInjectionFormat([
|
||||
{ content: 'no source' },
|
||||
])
|
||||
expect(text).toBe('<memory>\n<entry>no source</entry>\n</memory>')
|
||||
})
|
||||
|
||||
it('escapes XML in entry content and source so untrusted text cannot break the block', () => {
|
||||
const mm = new MemoryManager({ stores: [createMockStore('s')], injection: true })
|
||||
const text = (mm as unknown as { _defaultInjectionFormat(e: MemoryEntry[]): string })._defaultInjectionFormat([
|
||||
{ content: 'a < b & c > d </entry>', storeName: 'pre"f' },
|
||||
])
|
||||
expect(text).toBe(
|
||||
'<memory>\n<entry source="pre"f">a < b & c > d </entry></entry>\n</memory>'
|
||||
)
|
||||
})
|
||||
|
||||
it('honors a custom format', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'A' }] })
|
||||
const mm = new MemoryManager({
|
||||
stores: [store],
|
||||
injection: { format: ({ entries }) => `[${entries.map((e) => e.content).join('|')}]` },
|
||||
})
|
||||
|
||||
await expect(provide(mm, [assistant('prior'), user('ask')])).resolves.toBe('[A]')
|
||||
})
|
||||
|
||||
it('fails open (returns undefined) when a custom format throws', async () => {
|
||||
const store = createMockStore('s', { entries: [{ content: 'fact' }] })
|
||||
const mm = new MemoryManager({
|
||||
stores: [store],
|
||||
injection: {
|
||||
format: () => {
|
||||
throw new Error('boom')
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
// search ran, but the format threw, so nothing is injected.
|
||||
await expect(provide(mm, [assistant('prior'), user('ask')])).resolves.toBeUndefined()
|
||||
expect(store.search).toHaveBeenCalled()
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe('AgentConfig integration', () => {
|
||||
|
||||
@@ -10,7 +10,9 @@ export type {
|
||||
MemoryToolConfig,
|
||||
MemoryAddToolConfig,
|
||||
MemoryManagerConfig,
|
||||
MemoryInjectionConfig,
|
||||
} from './types.js'
|
||||
export type { InjectionConfig, InjectionTrigger, InjectionContext } from '../injection/index.js'
|
||||
|
||||
export { ExtractionTrigger } from './extraction/types.js'
|
||||
export { InvocationTrigger, IntervalTrigger } from './extraction/triggers.js'
|
||||
|
||||
@@ -3,6 +3,7 @@ import type { LocalAgent } from '../types/agent.js'
|
||||
import type { Tool } from '../tools/tool.js'
|
||||
import type {
|
||||
MemoryEntry,
|
||||
MemoryInjectionConfig,
|
||||
MemoryManagerConfig,
|
||||
MemorySearchOptions,
|
||||
MemoryStore,
|
||||
@@ -11,6 +12,7 @@ import type {
|
||||
MemoryAddToolConfig,
|
||||
} from './types.js'
|
||||
import type { JSONValue } from '../types/json.js'
|
||||
import type { MessageData } from '../types/messages.js'
|
||||
import { MessageAddedEvent } from '../hooks/events.js'
|
||||
import { ExtractionCoordinator, type ExtractionBinding } from './extraction/coordinator.js'
|
||||
import { resolveExtractionConfig } from './extraction/resolve-extraction-config.js'
|
||||
@@ -18,6 +20,9 @@ import { tool } from '../tools/tool-factory.js'
|
||||
import { z } from 'zod'
|
||||
import { logger } from '../logging/logger.js'
|
||||
import { normalizeError } from '../errors.js'
|
||||
import { isUserTurn, createInjectionMiddleware } from '../injection/message-injection.js'
|
||||
import { escapeXmlText, escapeXmlAttr } from '../injection/xml.js'
|
||||
import { InvokeModelStage } from '../middleware/index.js'
|
||||
|
||||
const SEARCH_TOOL_DESCRIPTION =
|
||||
'Search long-term memory for facts, preferences, or context from previous conversations. Use when you need background about the user or topic that may have been discussed before.'
|
||||
@@ -31,6 +36,16 @@ const ADD_TOOL_DESCRIPTION =
|
||||
*/
|
||||
export const DEFAULT_MAX_SEARCH_RESULTS = 3
|
||||
|
||||
/**
|
||||
* Default number of entries injected per model call when injection does not specify one.
|
||||
*
|
||||
* A memory store ranks by semantic (embedding) similarity, which is not the same as contextual
|
||||
* usefulness — the top hit is not reliably the most useful entry for the turn. Injecting the top few
|
||||
* gives the model a small candidate set to pick from rather than betting on the store's first result.
|
||||
* Five balances that recall against context bloat; lower it for a tighter prepend.
|
||||
*/
|
||||
const DEFAULT_MAX_ENTRIES = 5
|
||||
|
||||
/** Flattens nested AggregateErrors so the leaves are concrete reasons, not errors-of-errors. */
|
||||
function _flattenReasons(reasons: unknown[]): unknown[] {
|
||||
return reasons.flatMap((reason) => (reason instanceof AggregateError ? _flattenReasons(reason.errors) : [reason]))
|
||||
@@ -80,6 +95,8 @@ export class MemoryManager implements Plugin {
|
||||
private readonly _extractionStores: ExtractionBinding[]
|
||||
/** Background extraction coordinator, created in {@link initAgent} when extraction is configured. */
|
||||
private _coordinator: ExtractionCoordinator | undefined
|
||||
/** Resolved injection config, or `false` when injection is disabled. */
|
||||
private readonly _injectionConfig: MemoryInjectionConfig | false
|
||||
|
||||
constructor(config: MemoryManagerConfig) {
|
||||
if (config.stores.length === 0) {
|
||||
@@ -147,6 +164,13 @@ export class MemoryManager implements Plugin {
|
||||
this._addToolConfig = typeof config.addToolConfig === 'object' ? config.addToolConfig : {}
|
||||
this._addToolStores = this._resolveAddToolStores(this._addToolConfig)
|
||||
}
|
||||
|
||||
this._injectionConfig =
|
||||
config.injection === undefined || config.injection === false
|
||||
? false
|
||||
: typeof config.injection === 'object'
|
||||
? config.injection
|
||||
: {}
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -179,12 +203,22 @@ export class MemoryManager implements Plugin {
|
||||
/**
|
||||
* Initializes the plugin with the agent.
|
||||
*
|
||||
* Wires up automatic extraction for any store configured with {@link ExtractionConfig}: buffers
|
||||
* conversation messages and attaches each store's triggers. A no-op when no store uses extraction.
|
||||
* Wires up two independent behaviors:
|
||||
* - **Extraction**: for any store configured with {@link ExtractionConfig}, buffers conversation
|
||||
* messages and attaches each store's triggers. A no-op when no store uses extraction.
|
||||
* - **Injection**: when enabled, registers an `InvokeModelStage` middleware that folds retrieved
|
||||
* memory into the model input for each call without touching durable history. See
|
||||
* {@link _provideMemoryContext}, the `renderContent` callback the middleware invokes.
|
||||
*
|
||||
* @param agent - The agent this plugin is being attached to
|
||||
*/
|
||||
initAgent(agent: LocalAgent): void {
|
||||
this._initExtraction(agent)
|
||||
this._initInjection(agent)
|
||||
}
|
||||
|
||||
/** Wires background extraction for stores configured with {@link ExtractionConfig}. */
|
||||
private _initExtraction(agent: LocalAgent): void {
|
||||
if (this._extractionStores.length === 0) {
|
||||
return
|
||||
}
|
||||
@@ -204,6 +238,60 @@ export class MemoryManager implements Plugin {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Registers the injection middleware when injection is enabled. Folds retrieved memory into the
|
||||
* model input for each call via {@link _provideMemoryContext}, without touching durable history. A
|
||||
* no-op when injection is disabled.
|
||||
*/
|
||||
private _initInjection(agent: LocalAgent): void {
|
||||
const config = this._injectionConfig
|
||||
if (config === false) {
|
||||
return
|
||||
}
|
||||
agent.addMiddleware(
|
||||
InvokeModelStage.Input,
|
||||
createInjectionMiddleware({
|
||||
...(config.trigger !== undefined && { trigger: config.trigger }),
|
||||
renderContent: (context) => this._provideMemoryContext(context.messages, config),
|
||||
})
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* Produces the memory context text to inject for a model call, or `undefined` to skip. This is the
|
||||
* `renderContent` callback the injection middleware invokes (see {@link initAgent}).
|
||||
*
|
||||
* Derives a query (the configured callback or an adaptive default), searches memory, and renders the
|
||||
* top entries. Skips silently (returns `undefined`) when no query can be derived or the search
|
||||
* returns nothing. The rendering callback throwing fails open (returns `undefined`).
|
||||
*
|
||||
* @param messages - The current conversation, as data
|
||||
* @param config - The resolved injection configuration
|
||||
* @returns The injected text, or `undefined` when there is nothing to inject
|
||||
*/
|
||||
private async _provideMemoryContext(
|
||||
messages: MessageData[],
|
||||
config: MemoryInjectionConfig
|
||||
): Promise<string | undefined> {
|
||||
const query = this._resolveInjectionQuery(messages, config)
|
||||
if (!query?.trim()) {
|
||||
return undefined
|
||||
}
|
||||
|
||||
const maxResults = config.maxEntries ?? DEFAULT_MAX_ENTRIES
|
||||
const entries = (await this.search(query, { maxSearchResults: maxResults })).slice(0, maxResults)
|
||||
if (entries.length === 0) {
|
||||
return undefined
|
||||
}
|
||||
|
||||
try {
|
||||
return config.format ? config.format({ entries }) : this._defaultInjectionFormat(entries)
|
||||
} catch (error) {
|
||||
logger.warn(`reason=<${normalizeError(error).message}> | injection format threw; skipping injection`)
|
||||
return undefined
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Saves every store's remaining messages and waits for all saves to finish. No-op when no store has
|
||||
* extraction configured.
|
||||
@@ -220,6 +308,63 @@ export class MemoryManager implements Plugin {
|
||||
await this._coordinator?.flush()
|
||||
}
|
||||
|
||||
/**
|
||||
* Derives the injection search query. Uses the configured `query` callback when provided (a throw
|
||||
* fails open, skipping injection); otherwise an adaptive default: the latest user message's text on
|
||||
* a user turn, or the most recent assistant message's text otherwise (the previous autonomous step).
|
||||
*
|
||||
* @param messages - The current conversation, as data
|
||||
* @param config - The resolved injection configuration
|
||||
* @returns The query string, or `undefined` when none is available
|
||||
*/
|
||||
private _resolveInjectionQuery(messages: MessageData[], config: MemoryInjectionConfig): string | undefined {
|
||||
if (config.query) {
|
||||
try {
|
||||
return config.query({ messages })
|
||||
} catch (error) {
|
||||
logger.warn(`reason=<${normalizeError(error).message}> | injection query threw; skipping injection`)
|
||||
return undefined
|
||||
}
|
||||
}
|
||||
|
||||
const role = isUserTurn(messages) ? 'user' : 'assistant'
|
||||
const index = this._findLastIndex(messages, (message) => message.role === role)
|
||||
if (index < 0) {
|
||||
return undefined
|
||||
}
|
||||
const text = messages[index]!.content.filter((block) => 'text' in block)
|
||||
.map((block) => (block as { text: string }).text)
|
||||
.join('\n')
|
||||
.trim()
|
||||
return text.length > 0 ? text : undefined
|
||||
}
|
||||
|
||||
/**
|
||||
* Default injection format: a `<memory>` block with one `<entry>` per result. Each entry carries a
|
||||
* `source` attribute naming the originating store (when known) so the model can attribute memories.
|
||||
*
|
||||
* @param entries - The retrieved memory entries to render
|
||||
* @returns The rendered `<memory>` block
|
||||
*/
|
||||
private _defaultInjectionFormat(entries: MemoryEntry[]): string {
|
||||
const items = entries.map((entry) =>
|
||||
entry.storeName
|
||||
? `<entry source="${escapeXmlAttr(entry.storeName)}">${escapeXmlText(entry.content)}</entry>`
|
||||
: `<entry>${escapeXmlText(entry.content)}</entry>`
|
||||
)
|
||||
return `<memory>\n${items.join('\n')}\n</memory>`
|
||||
}
|
||||
|
||||
/** Returns the index of the last element matching `predicate`, or -1. */
|
||||
private _findLastIndex<T>(items: T[], predicate: (item: T) => boolean): number {
|
||||
for (let i = items.length - 1; i >= 0; i--) {
|
||||
if (predicate(items[i]!)) {
|
||||
return i
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns tools registered by this plugin.
|
||||
*
|
||||
|
||||
@@ -2,6 +2,7 @@ import type { JSONValue } from '../types/json.js'
|
||||
import type { MessageData } from '../types/messages.js'
|
||||
import type { Tool } from '../tools/tool.js'
|
||||
import type { ExtractionConfig } from './extraction/types.js'
|
||||
import type { InjectionConfig } from '../injection/index.js'
|
||||
|
||||
/**
|
||||
* A single memory entry retrieved from or stored to a memory store.
|
||||
@@ -206,6 +207,51 @@ export interface MemoryAddToolConfig extends MemoryToolConfig {
|
||||
waitForWrites?: boolean
|
||||
}
|
||||
|
||||
/**
|
||||
* Configuration for memory context injection.
|
||||
*
|
||||
* When enabled on a {@link MemoryManager}, the manager searches memory before a model call and makes
|
||||
* the top results available to the model for that call, so relevant knowledge is present without the
|
||||
* model choosing to search. The injected text is ephemeral: it augments the model input for that call
|
||||
* only and never persists into the durable conversation or session.
|
||||
*
|
||||
* Extends the generic {@link InjectionConfig} (which carries `trigger`) with the memory-owned knobs:
|
||||
* how many entries to retrieve, how to derive the query, and how to render the results.
|
||||
*/
|
||||
export interface MemoryInjectionConfig extends InjectionConfig {
|
||||
/**
|
||||
* Maximum number of entries to retrieve and inject per model call.
|
||||
*
|
||||
* A store ranks by semantic similarity, which is not the same as contextual usefulness, so the
|
||||
* default injects a small candidate set rather than betting on the top hit. Raising this improves
|
||||
* recall at the cost of a larger prepend (context bloat); lower it for a tighter injection.
|
||||
*
|
||||
* With multiple stores, results are concatenated in store-registration order with no cross-store
|
||||
* ranking, so this cap can favor entries from earlier-registered stores.
|
||||
*
|
||||
* @defaultValue 5
|
||||
*/
|
||||
maxEntries?: number
|
||||
/**
|
||||
* Derives the search query from the current conversation. Return `undefined` or an empty string to
|
||||
* skip injection for this call. A callback that throws fails open (injection is skipped).
|
||||
*
|
||||
* Defaults to an adaptive query: the latest user message's text on a user turn, otherwise the most
|
||||
* recent assistant message's text (the previous step on an autonomous turn).
|
||||
*/
|
||||
query?: (context: { messages: MessageData[] }) => string | undefined
|
||||
/**
|
||||
* Renders retrieved entries into the injected text. A callback that throws fails open (injection is
|
||||
* skipped).
|
||||
*
|
||||
* Defaults to a `<memory>` XML block with one `<entry>` per result, carrying a `source` attribute
|
||||
* naming the originating store (when known) so the model can attribute and weigh each memory. The
|
||||
* default escapes entry content and source, so a custom `format` that emits markup is responsible
|
||||
* for its own escaping.
|
||||
*/
|
||||
format?: (context: { entries: MemoryEntry[] }) => string
|
||||
}
|
||||
|
||||
/**
|
||||
* Configuration for the {@link MemoryManager}.
|
||||
*/
|
||||
@@ -219,4 +265,19 @@ export interface MemoryManagerConfig {
|
||||
* writable stores; pass a {@link MemoryAddToolConfig} with `stores` to restrict it to specific ones.
|
||||
*/
|
||||
addToolConfig?: MemoryAddToolConfig | boolean
|
||||
/**
|
||||
* Memory context injection. Defaults to `false` (opt-in). `true` uses the default injection
|
||||
* settings; pass a {@link MemoryInjectionConfig} to customize retrieval, timing, and formatting.
|
||||
*
|
||||
* `true` is equivalent to:
|
||||
* ```ts
|
||||
* {
|
||||
* trigger: 'userTurn', // inject only on a fresh user ask
|
||||
* maxEntries: 5, // retrieve and inject up to 5 entries
|
||||
* // query: the latest user text on a user turn, else the most recent assistant text
|
||||
* // format: a <memory> block with one <entry source="STORE_NAME"> per result (content escaped)
|
||||
* }
|
||||
* ```
|
||||
*/
|
||||
injection?: boolean | MemoryInjectionConfig
|
||||
}
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
import { describe, it, expect, vi } from 'vitest'
|
||||
import { ContextInjector } from '../plugin.js'
|
||||
import { InvokeModelStage } from '../../../middleware/index.js'
|
||||
import { Message, TextBlock } from '../../../types/messages.js'
|
||||
import type { InvokeModelContext } from '../../../middleware/index.js'
|
||||
import { createMockAgent } from '../../../__fixtures__/agent-helpers.js'
|
||||
|
||||
const user = (text: string) => new Message({ role: 'user', content: [new TextBlock(text)] })
|
||||
const assistant = (text: string) => new Message({ role: 'assistant', content: [new TextBlock(text)] })
|
||||
|
||||
// Builds a mock agent that captures addMiddleware registrations, with a real appState/cancelSignal.
|
||||
function makeAgent() {
|
||||
const addMiddleware = vi.fn()
|
||||
const agent = createMockAgent({ extra: { addMiddleware } as never })
|
||||
return { agent, addMiddleware }
|
||||
}
|
||||
|
||||
describe('ContextInjector', () => {
|
||||
describe('plugin interface', () => {
|
||||
it('defaults to the strands:context-injector name', () => {
|
||||
expect(new ContextInjector({ renderContent: async () => 'x' }).name).toBe('strands:context-injector')
|
||||
})
|
||||
|
||||
it('honors a custom name (so multiple injectors can be told apart)', () => {
|
||||
expect(new ContextInjector({ name: 'now', renderContent: async () => 'x' }).name).toBe('now')
|
||||
})
|
||||
|
||||
it('registers an InvokeModelStage input middleware on initAgent', () => {
|
||||
const { agent, addMiddleware } = makeAgent()
|
||||
new ContextInjector({ renderContent: async () => 'x' }).initAgent(agent)
|
||||
|
||||
expect(addMiddleware).toHaveBeenCalledTimes(1)
|
||||
expect(addMiddleware.mock.calls[0]![0]).toBe(InvokeModelStage.Input)
|
||||
expect(typeof addMiddleware.mock.calls[0]![1]).toBe('function')
|
||||
})
|
||||
})
|
||||
|
||||
describe('registered handler', () => {
|
||||
// Runs the handler the plugin registered, with a context backed by the mock agent.
|
||||
async function run(plugin: ContextInjector, messages: Message[]) {
|
||||
const { agent, addMiddleware } = makeAgent()
|
||||
plugin.initAgent(agent)
|
||||
const handler = addMiddleware.mock.calls[0]![1] as (ctx: InvokeModelContext) => Promise<InvokeModelContext>
|
||||
return handler({ messages, agent } as unknown as InvokeModelContext)
|
||||
}
|
||||
|
||||
it('folds renderContent() text into the latest user message', async () => {
|
||||
const result = await run(new ContextInjector({ renderContent: async () => 'INJECTED' }), [
|
||||
assistant('prior'),
|
||||
user('ask'),
|
||||
])
|
||||
expect(result.messages.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'assistant', content: [{ text: 'prior' }] },
|
||||
{ role: 'user', content: [{ text: 'INJECTED' }, { text: 'ask' }] },
|
||||
])
|
||||
})
|
||||
|
||||
it('skips on a non-user turn by default (userTurn trigger)', async () => {
|
||||
const renderContent = vi.fn(async () => 'x')
|
||||
const input = [user('ask'), assistant('reply')]
|
||||
const result = await run(new ContextInjector({ renderContent }), input)
|
||||
expect(renderContent).not.toHaveBeenCalled()
|
||||
expect(result.messages).toBe(input)
|
||||
})
|
||||
|
||||
it("'everyTurn' injects regardless of the latest role", async () => {
|
||||
const result = await run(new ContextInjector({ trigger: 'everyTurn', renderContent: async () => 'INJECTED' }), [
|
||||
user('ask'),
|
||||
assistant('reply'),
|
||||
])
|
||||
// No later user message than index 0, so the fold targets it.
|
||||
expect(result.messages.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'user', content: [{ text: 'INJECTED' }, { text: 'ask' }] },
|
||||
{ role: 'assistant', content: [{ text: 'reply' }] },
|
||||
])
|
||||
})
|
||||
|
||||
it('exposes appState and the agent to renderContent', async () => {
|
||||
const { agent, addMiddleware } = makeAgent()
|
||||
let sawAgent = false
|
||||
let sawAppState = false
|
||||
new ContextInjector({
|
||||
renderContent: async (ctx) => {
|
||||
sawAgent = ctx.agent === agent
|
||||
sawAppState = ctx.appState === agent.appState
|
||||
return undefined
|
||||
},
|
||||
}).initAgent(agent)
|
||||
const handler = addMiddleware.mock.calls[0]![1] as (ctx: InvokeModelContext) => Promise<InvokeModelContext>
|
||||
await handler({ messages: [user('ask')], agent } as unknown as InvokeModelContext)
|
||||
|
||||
expect(sawAgent).toBe(true)
|
||||
expect(sawAppState).toBe(true)
|
||||
})
|
||||
|
||||
it('fails open (passes context through) when renderContent throws', async () => {
|
||||
const result = await run(
|
||||
new ContextInjector({
|
||||
renderContent: async () => {
|
||||
throw new Error('boom')
|
||||
},
|
||||
}),
|
||||
[assistant('prior'), user('ask')]
|
||||
)
|
||||
// Unchanged: the original messages, no injected block.
|
||||
expect(result.messages.map((m) => m.toJSON())).toStrictEqual([
|
||||
{ role: 'assistant', content: [{ text: 'prior' }] },
|
||||
{ role: 'user', content: [{ text: 'ask' }] },
|
||||
])
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,25 @@
|
||||
/**
|
||||
* Context-injection plugin for Strands Agents.
|
||||
*
|
||||
* Provides the {@link ContextInjector} plugin, which folds just-in-time text into the model input
|
||||
* before each call without touching durable history.
|
||||
*
|
||||
* @example
|
||||
* ```typescript
|
||||
* import { Agent } from '@strands-agents/sdk'
|
||||
* import { ContextInjector } from '@strands-agents/sdk/vended-plugins/context-injector'
|
||||
*
|
||||
* const agent = new Agent({
|
||||
* model,
|
||||
* plugins: [
|
||||
* new ContextInjector({
|
||||
* renderContent: async ({ messages }) => `<context>${derive(messages)}</context>`,
|
||||
* }),
|
||||
* ],
|
||||
* })
|
||||
* ```
|
||||
*/
|
||||
|
||||
export { ContextInjector } from './plugin.js'
|
||||
export type { ContextInjectorConfig } from './plugin.js'
|
||||
export type { InjectionTrigger, InjectionContext } from '../../injection/types.js'
|
||||
@@ -0,0 +1,73 @@
|
||||
import type { Plugin } from '../../plugins/plugin.js'
|
||||
import type { LocalAgent } from '../../types/agent.js'
|
||||
import { InvokeModelStage } from '../../middleware/index.js'
|
||||
import { createInjectionMiddleware } from '../../injection/message-injection.js'
|
||||
import type { InjectionTrigger, InjectionContext } from '../../injection/types.js'
|
||||
|
||||
/** Configuration for the {@link ContextInjector} plugin. */
|
||||
export interface ContextInjectorConfig {
|
||||
/**
|
||||
* Plugin name, for logging and duplicate detection. Defaults to `'strands:context-injector'`. Set a
|
||||
* distinct name when registering more than one injector so they can be told apart.
|
||||
*/
|
||||
name?: string
|
||||
/**
|
||||
* When to inject. An {@link InjectionTrigger} name selects a built-in policy (`'userTurn'` —
|
||||
* default — or `'everyTurn'`); a predicate over the {@link InjectionContext} is the escape hatch. A
|
||||
* predicate that throws fails open (injection is skipped).
|
||||
*
|
||||
* @defaultValue 'userTurn'
|
||||
*/
|
||||
trigger?: InjectionTrigger | ((context: InjectionContext) => boolean)
|
||||
/**
|
||||
* Renders the text to inject for this call, or `undefined`/`''` to skip. The text reaches the model
|
||||
* verbatim, so it is a prompt-injection surface: escape any attacker-influenced fields yourself. A
|
||||
* callback that throws fails open (injection is skipped, the model call proceeds).
|
||||
*/
|
||||
renderContent: (context: InjectionContext) => Promise<string | undefined>
|
||||
}
|
||||
|
||||
/**
|
||||
* Plugin that injects just-in-time context into the model input before each call.
|
||||
*
|
||||
* Before each model call, the plugin asks {@link ContextInjectorConfig.renderContent} for text and
|
||||
* makes it available to the model for that call, gated by {@link ContextInjectorConfig.trigger}. The
|
||||
* injected text is ephemeral: it augments the model input for that one call and never persists into the
|
||||
* durable conversation or session.
|
||||
*
|
||||
* @example
|
||||
* ```typescript
|
||||
* import { Agent } from '@strands-agents/sdk'
|
||||
* import { ContextInjector } from '@strands-agents/sdk/vended-plugins/context-injector'
|
||||
*
|
||||
* const agent = new Agent({
|
||||
* model,
|
||||
* plugins: [new ContextInjector({ renderContent: async () => `<now>${new Date().toISOString()}</now>` })],
|
||||
* })
|
||||
* ```
|
||||
*
|
||||
* @remarks
|
||||
* Multiple injectors may be registered; each contributes its text independently, in plugin-registration
|
||||
* order.
|
||||
*/
|
||||
export class ContextInjector implements Plugin {
|
||||
readonly name: string
|
||||
|
||||
private readonly _config: ContextInjectorConfig
|
||||
|
||||
constructor(config: ContextInjectorConfig) {
|
||||
this.name = config.name ?? 'strands:context-injector'
|
||||
this._config = config
|
||||
}
|
||||
|
||||
initAgent(agent: LocalAgent): void {
|
||||
const config = this._config
|
||||
agent.addMiddleware(
|
||||
InvokeModelStage.Input,
|
||||
createInjectionMiddleware({
|
||||
...(config.trigger !== undefined && { trigger: config.trigger }),
|
||||
renderContent: config.renderContent,
|
||||
})
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -3,10 +3,11 @@
|
||||
*
|
||||
* Provides a single import path for consumers who want all built-in plugins:
|
||||
* ```typescript
|
||||
* import { AgentSkills, ContextOffloader, GoalLoop, InMemoryStorage } from '@strands-agents/sdk/vended-plugins'
|
||||
* import { AgentSkills, ContextOffloader, ContextInjector, GoalLoop, InMemoryStorage } from '@strands-agents/sdk/vended-plugins'
|
||||
* ```
|
||||
*/
|
||||
|
||||
export * from './skills/index.js'
|
||||
export * from './context-offloader/index.js'
|
||||
export * from './context-injector/index.js'
|
||||
export * from './goal/index.js'
|
||||
|
||||
Reference in New Issue
Block a user