@empyria/restate 0.1.23 → 0.2.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/AGENTS.md +12 -0
- package/README.md +60 -0
- package/agent.js +11 -0
- package/lib/agent/Agent.js +246 -0
- package/lib/agent/Errors.js +86 -0
- package/lib/agent/Loop.js +157 -0
- package/lib/agent/Model.js +19 -0
- package/lib/agent/Skills.js +81 -0
- package/lib/agent/Tools.js +215 -0
- package/package.json +21 -1
- package/test/agent/Agent.test.js +195 -0
- package/test/agent/Errors.test.js +66 -0
- package/test/agent/Loop.test.js +173 -0
- package/test/agent/Skills.test.js +147 -0
- package/test/agent/Tools.test.js +277 -0
- package/test/agent/fakes.js +84 -0
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
import { createHash } from 'node:crypto'
|
|
2
|
+
import { TerminalError } from '@restatedev/restate-sdk'
|
|
3
|
+
|
|
4
|
+
export const LOAD_SKILL_TOOL = 'load_skill'
|
|
5
|
+
|
|
6
|
+
/**
|
|
7
|
+
* @typedef {{name: string, sha256: string}} SkillRef What the journal records for a loaded
|
|
8
|
+
* skill instead of its body.
|
|
9
|
+
*/
|
|
10
|
+
|
|
11
|
+
/** @param {string} body */
|
|
12
|
+
export const hashSkill = (body) => createHash('sha256').update(body).digest('hex')
|
|
13
|
+
|
|
14
|
+
/**
|
|
15
|
+
* The `load_skill` tool result as it's journaled and kept in history: a reference only.
|
|
16
|
+
* {@link expandSkillRefs} turns it back into the body right before each model call.
|
|
17
|
+
* @param {SkillRef} skillRef
|
|
18
|
+
*/
|
|
19
|
+
export const skillRefOutput = (skillRef) => ({ type: 'json', value: { skillRef } })
|
|
20
|
+
|
|
21
|
+
/**
|
|
22
|
+
* Returns a copy of `messages` with every `load_skill` reference replaced by the skill's
|
|
23
|
+
* body. Called inside each LLM step's `ctx.run`, whose input is never journaled — so skill
|
|
24
|
+
* bodies never reach the journal, only the {@link SkillRef}s do.
|
|
25
|
+
*
|
|
26
|
+
* `strict` is for references recorded in the current invocation: the body must hash to the
|
|
27
|
+
* recorded value, otherwise the skill changed between the original execution and a replay,
|
|
28
|
+
* and the model would silently see different instructions than before — a `TerminalError`.
|
|
29
|
+
* Non-strict is for references from earlier turns of a session: a skill updated between
|
|
30
|
+
* turns is legitimate, so the current body is used, or a notice if the skill is gone.
|
|
31
|
+
* @param {Array<object>} messages
|
|
32
|
+
* @param {{fetchSkill?: (name: string) => Promise<string|undefined>, strict: boolean}} params
|
|
33
|
+
* @returns {Promise<Array<object>>}
|
|
34
|
+
*/
|
|
35
|
+
export async function expandSkillRefs(messages, { fetchSkill, strict }) {
|
|
36
|
+
return Promise.all(
|
|
37
|
+
messages.map(async (message) => {
|
|
38
|
+
if (message.role !== 'tool' || !Array.isArray(message.content)) return message
|
|
39
|
+
return {
|
|
40
|
+
...message,
|
|
41
|
+
content: await Promise.all(
|
|
42
|
+
message.content.map((part) => expandPart(part, { fetchSkill, strict })),
|
|
43
|
+
),
|
|
44
|
+
}
|
|
45
|
+
}),
|
|
46
|
+
)
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
async function expandPart(part, { fetchSkill, strict }) {
|
|
50
|
+
const skillRef = part.type === 'tool-result' && part.output?.value?.skillRef
|
|
51
|
+
if (part.toolName !== LOAD_SKILL_TOOL || !skillRef) return part
|
|
52
|
+
|
|
53
|
+
const body = (await fetchSkill?.(skillRef.name)) ?? undefined
|
|
54
|
+
if (body !== undefined && (!strict || hashSkill(body) === skillRef.sha256)) {
|
|
55
|
+
return { ...part, output: { type: 'text', value: body } }
|
|
56
|
+
}
|
|
57
|
+
if (strict) {
|
|
58
|
+
throw new TerminalError(
|
|
59
|
+
`Skill "${skillRef.name}" changed or disappeared during this invocation; refusing to send the model different instructions than it saw before`,
|
|
60
|
+
)
|
|
61
|
+
}
|
|
62
|
+
return {
|
|
63
|
+
...part,
|
|
64
|
+
output: { type: 'error-text', value: `Skill "${skillRef.name}" is no longer available.` },
|
|
65
|
+
}
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
/**
|
|
69
|
+
* Memoises `loadSkill` for one execution of the handler. A replay starts a new execution
|
|
70
|
+
* and so a fresh cache; the hash check in {@link expandSkillRefs} catches content that
|
|
71
|
+
* changed in between.
|
|
72
|
+
* @param {(name: string) => Promise<string|undefined>} [loadSkill]
|
|
73
|
+
*/
|
|
74
|
+
export function cachedSkillFetcher(loadSkill) {
|
|
75
|
+
if (!loadSkill) return undefined
|
|
76
|
+
const cache = new Map()
|
|
77
|
+
return (name) => {
|
|
78
|
+
if (!cache.has(name)) cache.set(name, Promise.resolve(loadSkill(name)))
|
|
79
|
+
return cache.get(name)
|
|
80
|
+
}
|
|
81
|
+
}
|
|
@@ -0,0 +1,215 @@
|
|
|
1
|
+
import { tool, jsonSchema } from 'ai'
|
|
2
|
+
import { CancelledError, TerminalError, TimeoutError, serde } from '@restatedev/restate-sdk'
|
|
3
|
+
import { defineSchema, string } from '@empyria/common'
|
|
4
|
+
import { DEFAULT_TOOL_RETRY } from './Errors.js'
|
|
5
|
+
import { LOAD_SKILL_TOOL, hashSkill, skillRefOutput } from './Skills.js'
|
|
6
|
+
|
|
7
|
+
/**
|
|
8
|
+
* Tool names the agent loop owns itself; a user-declared tool can't reuse them.
|
|
9
|
+
*/
|
|
10
|
+
export { LOAD_SKILL_TOOL }
|
|
11
|
+
export const FINAL_ANSWER_TOOL = 'final_answer'
|
|
12
|
+
|
|
13
|
+
const TOOL_NAME = /^[a-zA-Z0-9_-]{1,64}$/
|
|
14
|
+
|
|
15
|
+
/**
|
|
16
|
+
* @typedef {Object} ApprovalSpec Human (or system) approval required before a tool runs.
|
|
17
|
+
* @property {number} timeout Milliseconds to wait for a decision. A timeout counts as a
|
|
18
|
+
* rejection.
|
|
19
|
+
* @property {{block: string|{name: string}, handler: string, key?: string}} notify Handler
|
|
20
|
+
* that's sent `{awakeableId, tool, input, invocationId}` when approval is needed. The
|
|
21
|
+
* approver answers by resolving the awakeable with `{approved: boolean, reason?: string}`
|
|
22
|
+
* (Restate ingress: `POST /restate/awakeables/<awakeableId>/resolve`).
|
|
23
|
+
*
|
|
24
|
+
* @typedef {Object} BlockToolSpec A tool executed as a Restate call to another building block.
|
|
25
|
+
* Restate journals the call, retries it inside the callee and guarantees it runs once.
|
|
26
|
+
* @property {string} description
|
|
27
|
+
* @property {object} inputSchema JSON Schema of the tool's arguments, shown to the model.
|
|
28
|
+
* @property {string|{name: string}} block Target service/object/workflow, by name or definition.
|
|
29
|
+
* @property {string} handler Target handler.
|
|
30
|
+
* @property {(args: any) => string} [key] Virtual object / workflow key, derived from the arguments.
|
|
31
|
+
* @property {ApprovalSpec} [approval]
|
|
32
|
+
*
|
|
33
|
+
* @typedef {Object} ExecuteToolSpec A tool executed as a direct side effect inside its own `ctx.run`.
|
|
34
|
+
* @property {string} description
|
|
35
|
+
* @property {object} inputSchema JSON Schema of the tool's arguments, shown to the model.
|
|
36
|
+
* @property {(args: any, meta: {idempotencyKey: string}) => any} execute Runs the side effect.
|
|
37
|
+
* May run more than once before its result is recorded: pass `idempotencyKey` on to the
|
|
38
|
+
* external system for any write. Must return a small JSON-serialisable value.
|
|
39
|
+
* @property {import('@restatedev/restate-sdk').RunOptions<any>} [retry] Defaults to
|
|
40
|
+
* {@link DEFAULT_TOOL_RETRY}.
|
|
41
|
+
* @property {ApprovalSpec} [approval]
|
|
42
|
+
*
|
|
43
|
+
* @typedef {BlockToolSpec|ExecuteToolSpec} ToolSpec
|
|
44
|
+
*/
|
|
45
|
+
|
|
46
|
+
/**
|
|
47
|
+
* Validates a `tools` map at definition time, so a misconfigured agent fails on startup
|
|
48
|
+
* rather than on its first invocation.
|
|
49
|
+
* @param {Record<string, ToolSpec>} tools
|
|
50
|
+
* @throws {TypeError}
|
|
51
|
+
*/
|
|
52
|
+
export function assertToolSpecs(tools) {
|
|
53
|
+
for (const [name, spec] of Object.entries(tools)) {
|
|
54
|
+
if (!TOOL_NAME.test(name)) {
|
|
55
|
+
throw new TypeError(`Invalid tool name '${name}': use [a-zA-Z0-9_-], at most 64 chars`)
|
|
56
|
+
}
|
|
57
|
+
if (name === LOAD_SKILL_TOOL || name === FINAL_ANSWER_TOOL) {
|
|
58
|
+
throw new TypeError(`Tool name '${name}' is reserved by the agent loop`)
|
|
59
|
+
}
|
|
60
|
+
if (typeof spec.description !== 'string' || !spec.inputSchema) {
|
|
61
|
+
throw new TypeError(`Tool '${name}' needs a description and an inputSchema`)
|
|
62
|
+
}
|
|
63
|
+
const isBlock = spec.block !== undefined
|
|
64
|
+
const isExecute = typeof spec.execute === 'function'
|
|
65
|
+
if (isBlock === isExecute) {
|
|
66
|
+
throw new TypeError(`Tool '${name}' needs exactly one of 'block' or 'execute'`)
|
|
67
|
+
}
|
|
68
|
+
if (isBlock && typeof spec.handler !== 'string') {
|
|
69
|
+
throw new TypeError(`Block tool '${name}' needs a 'handler'`)
|
|
70
|
+
}
|
|
71
|
+
if (spec.approval) {
|
|
72
|
+
const { timeout, notify } = spec.approval
|
|
73
|
+
if (!(timeout > 0) || !notify?.block || !notify?.handler) {
|
|
74
|
+
throw new TypeError(
|
|
75
|
+
`Tool '${name}' approval needs a positive 'timeout' and a 'notify' {block, handler}`,
|
|
76
|
+
)
|
|
77
|
+
}
|
|
78
|
+
}
|
|
79
|
+
}
|
|
80
|
+
}
|
|
81
|
+
|
|
82
|
+
/**
|
|
83
|
+
* AI SDK tool definitions for the model. None has an `execute`: AI SDK then returns the
|
|
84
|
+
* requested call instead of running it, so {@link executeTool} owns every execution as
|
|
85
|
+
* its own durable step.
|
|
86
|
+
* @param {{tools: Record<string, ToolSpec>, hasSkills: boolean, outputSchema?: object}} params
|
|
87
|
+
*/
|
|
88
|
+
export function buildToolDefs({ tools, hasSkills, outputSchema }) {
|
|
89
|
+
const defs = Object.fromEntries(
|
|
90
|
+
Object.entries(tools).map(([name, spec]) => [
|
|
91
|
+
name,
|
|
92
|
+
tool({ description: spec.description, inputSchema: jsonSchema(spec.inputSchema) }),
|
|
93
|
+
]),
|
|
94
|
+
)
|
|
95
|
+
if (hasSkills) {
|
|
96
|
+
defs[LOAD_SKILL_TOOL] = tool({
|
|
97
|
+
description:
|
|
98
|
+
"Load the full instructions for a named skill, by the name shown in the system prompt's skill index.",
|
|
99
|
+
inputSchema: jsonSchema(defineSchema({ name: string() })),
|
|
100
|
+
})
|
|
101
|
+
}
|
|
102
|
+
if (outputSchema) {
|
|
103
|
+
defs[FINAL_ANSWER_TOOL] = tool({
|
|
104
|
+
description: 'Return your final answer. Call this exactly once, when you are done.',
|
|
105
|
+
inputSchema: jsonSchema(outputSchema),
|
|
106
|
+
})
|
|
107
|
+
}
|
|
108
|
+
return defs
|
|
109
|
+
}
|
|
110
|
+
|
|
111
|
+
/**
|
|
112
|
+
* Executes one model-requested tool call and returns the AI SDK tool-result `output`.
|
|
113
|
+
*
|
|
114
|
+
* Error policy (see the workflow model, §7): a tool that fails terminally — rejected
|
|
115
|
+
* approval, a `TerminalError` from the callee, an `execute` tool whose retries ran out —
|
|
116
|
+
* is reported back to the model as `error-text`, so the agent can adapt. Transient errors
|
|
117
|
+
* are never swallowed: Restate retries them. Cancellation always propagates.
|
|
118
|
+
* @param {import('@restatedev/restate-sdk').Context} ctx
|
|
119
|
+
* @param {{toolCallId: string, toolName: string, input: any}} call
|
|
120
|
+
* @param {{tools: Record<string, ToolSpec>, loadSkill?: (name: string) => Promise<string|undefined>, step: number, index: number}} params
|
|
121
|
+
*/
|
|
122
|
+
export async function executeTool(ctx, call, { tools, loadSkill, step, index }) {
|
|
123
|
+
const stepId = `${step}.${index}`
|
|
124
|
+
try {
|
|
125
|
+
if (call.toolName === LOAD_SKILL_TOOL && loadSkill) {
|
|
126
|
+
// Journals only the reference; the body is put back into the conversation inside
|
|
127
|
+
// each LLM step, never recorded (see Skills.js).
|
|
128
|
+
const skillRef = await ctx.run(`skill:${call.input?.name}-${stepId}`, async () => {
|
|
129
|
+
const body = await loadSkill(call.input?.name)
|
|
130
|
+
return typeof body === 'string'
|
|
131
|
+
? { name: call.input.name, sha256: hashSkill(body) }
|
|
132
|
+
: null
|
|
133
|
+
})
|
|
134
|
+
return skillRef === null
|
|
135
|
+
? {
|
|
136
|
+
type: 'error-text',
|
|
137
|
+
value: `No skill named "${call.input?.name}" is available.`,
|
|
138
|
+
}
|
|
139
|
+
: skillRefOutput(skillRef)
|
|
140
|
+
}
|
|
141
|
+
|
|
142
|
+
const spec = tools[call.toolName]
|
|
143
|
+
if (!spec) return { type: 'error-text', value: `Unknown tool "${call.toolName}".` }
|
|
144
|
+
|
|
145
|
+
if (spec.approval) {
|
|
146
|
+
const decision = await requestApproval(ctx, call, spec.approval)
|
|
147
|
+
if (!decision.approved) {
|
|
148
|
+
return {
|
|
149
|
+
type: 'error-text',
|
|
150
|
+
value: `Tool "${call.toolName}" was not approved: ${decision.reason ?? 'rejected'}.`,
|
|
151
|
+
}
|
|
152
|
+
}
|
|
153
|
+
}
|
|
154
|
+
|
|
155
|
+
const value = spec.execute
|
|
156
|
+
? await ctx.run(
|
|
157
|
+
`tool:${call.toolName}-${stepId}`,
|
|
158
|
+
async () =>
|
|
159
|
+
(await spec.execute(call.input, {
|
|
160
|
+
idempotencyKey: `${ctx.request().id}:${stepId}`,
|
|
161
|
+
})) ?? null,
|
|
162
|
+
spec.retry ?? DEFAULT_TOOL_RETRY,
|
|
163
|
+
)
|
|
164
|
+
: await ctx.genericCall({
|
|
165
|
+
service: targetName(spec.block),
|
|
166
|
+
method: spec.handler,
|
|
167
|
+
key: spec.key?.(call.input),
|
|
168
|
+
parameter: call.input,
|
|
169
|
+
inputSerde: serde.json,
|
|
170
|
+
outputSerde: serde.json,
|
|
171
|
+
})
|
|
172
|
+
|
|
173
|
+
return { type: 'json', value: value ?? null }
|
|
174
|
+
} catch (error) {
|
|
175
|
+
if (error instanceof CancelledError || !(error instanceof TerminalError)) throw error
|
|
176
|
+
return { type: 'error-text', value: error.message }
|
|
177
|
+
}
|
|
178
|
+
}
|
|
179
|
+
|
|
180
|
+
/**
|
|
181
|
+
* Asks for approval through an awakeable and waits for the decision, durably.
|
|
182
|
+
* @param {import('@restatedev/restate-sdk').Context} ctx
|
|
183
|
+
* @param {{toolName: string, input: any}} call
|
|
184
|
+
* @param {ApprovalSpec} approval
|
|
185
|
+
* @returns {Promise<{approved: boolean, reason?: string}>}
|
|
186
|
+
*/
|
|
187
|
+
async function requestApproval(ctx, call, { timeout, notify }) {
|
|
188
|
+
const { id, promise } = ctx.awakeable()
|
|
189
|
+
ctx.genericSend({
|
|
190
|
+
service: targetName(notify.block),
|
|
191
|
+
method: notify.handler,
|
|
192
|
+
key: notify.key,
|
|
193
|
+
parameter: {
|
|
194
|
+
awakeableId: id,
|
|
195
|
+
tool: call.toolName,
|
|
196
|
+
input: call.input,
|
|
197
|
+
invocationId: ctx.request().id,
|
|
198
|
+
},
|
|
199
|
+
inputSerde: serde.json,
|
|
200
|
+
})
|
|
201
|
+
|
|
202
|
+
let decision
|
|
203
|
+
try {
|
|
204
|
+
decision = await promise.orTimeout(timeout)
|
|
205
|
+
} catch (error) {
|
|
206
|
+
if (error instanceof TimeoutError) return { approved: false, reason: 'approval timed out' }
|
|
207
|
+
throw error
|
|
208
|
+
}
|
|
209
|
+
return {
|
|
210
|
+
approved: decision?.approved === true,
|
|
211
|
+
reason: typeof decision?.reason === 'string' ? decision.reason : undefined,
|
|
212
|
+
}
|
|
213
|
+
}
|
|
214
|
+
|
|
215
|
+
const targetName = (block) => (typeof block === 'string' ? block : block.name)
|
package/package.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
{
|
|
2
2
|
"name": "@empyria/restate",
|
|
3
|
-
"version": "0.
|
|
3
|
+
"version": "0.2.0",
|
|
4
4
|
"description": "Restate.dev helpers for the Empyria nanoservice framework",
|
|
5
5
|
"license": "MIT",
|
|
6
6
|
"author": "Imre Fazekas <imre.fazekas@icloud.com>",
|
|
@@ -12,6 +12,7 @@
|
|
|
12
12
|
"main": "./index.js",
|
|
13
13
|
"exports": {
|
|
14
14
|
".": "./index.js",
|
|
15
|
+
"./agent": "./agent.js",
|
|
15
16
|
"./lib/*": "./lib/*.js",
|
|
16
17
|
"./package.json": "./package.json"
|
|
17
18
|
},
|
|
@@ -34,10 +35,29 @@
|
|
|
34
35
|
"croner": "10.0.1"
|
|
35
36
|
},
|
|
36
37
|
"devDependencies": {
|
|
38
|
+
"@ai-sdk/openai-compatible": "3.0.60",
|
|
39
|
+
"@ai-sdk/provider": "4.0.20",
|
|
37
40
|
"@restatedev/restate-server": "1.7.12",
|
|
41
|
+
"ai": "7.0.123",
|
|
38
42
|
"oxfmt": "0.71.0",
|
|
39
43
|
"oxlint": "1.86.0"
|
|
40
44
|
},
|
|
45
|
+
"peerDependencies": {
|
|
46
|
+
"@ai-sdk/openai-compatible": "^3.0.60",
|
|
47
|
+
"@ai-sdk/provider": "^4.0.20",
|
|
48
|
+
"ai": "^7.0.123"
|
|
49
|
+
},
|
|
50
|
+
"peerDependenciesMeta": {
|
|
51
|
+
"@ai-sdk/openai-compatible": {
|
|
52
|
+
"optional": true
|
|
53
|
+
},
|
|
54
|
+
"@ai-sdk/provider": {
|
|
55
|
+
"optional": true
|
|
56
|
+
},
|
|
57
|
+
"ai": {
|
|
58
|
+
"optional": true
|
|
59
|
+
}
|
|
60
|
+
},
|
|
41
61
|
"engines": {
|
|
42
62
|
"bun": ">=1.4.2",
|
|
43
63
|
"node": ">=26"
|
|
@@ -0,0 +1,195 @@
|
|
|
1
|
+
import { describe, test, expect } from 'bun:test'
|
|
2
|
+
import { MockLanguageModelV4 } from 'ai/test'
|
|
3
|
+
import { TerminalError } from '@restatedev/restate-sdk'
|
|
4
|
+
import { createAgentService, defineAgent } from '../../lib/agent/Agent.js'
|
|
5
|
+
import { DEFAULT_AGENT_RETRY_POLICY } from '../../lib/agent/Errors.js'
|
|
6
|
+
import { env, fakeCtx, sequence, textResult, toolCallResult } from './fakes.js'
|
|
7
|
+
|
|
8
|
+
const okModel = () => new MockLanguageModelV4({ doGenerate: textResult('ok') })
|
|
9
|
+
|
|
10
|
+
describe('defineAgent: definition', () => {
|
|
11
|
+
test('memory none → a service named after the agent, with the default retry policy', () => {
|
|
12
|
+
const agent = defineAgent({ name: 'Triage', model: okModel() })
|
|
13
|
+
expect(agent.name).toBe('Triage')
|
|
14
|
+
expect(typeof agent.service.ask).toBe('function')
|
|
15
|
+
expect(agent.options.retryPolicy).toEqual(DEFAULT_AGENT_RETRY_POLICY)
|
|
16
|
+
})
|
|
17
|
+
|
|
18
|
+
test('memory session → a virtual object with ask and reset', () => {
|
|
19
|
+
const agent = defineAgent({ name: 'Chat', model: okModel(), memory: 'session' })
|
|
20
|
+
expect(Object.keys(agent.object).sort()).toEqual(['ask', 'reset'])
|
|
21
|
+
})
|
|
22
|
+
|
|
23
|
+
test.each([
|
|
24
|
+
['no name', { model: {} }],
|
|
25
|
+
['no model or env', { name: 'A' }],
|
|
26
|
+
['an unknown memory mode', { name: 'A', model: {}, memory: 'forever' }],
|
|
27
|
+
['a historyStore without session memory', { name: 'A', model: {}, historyStore: {} }],
|
|
28
|
+
['an invalid tool', { name: 'A', model: {}, tools: { load_skill: {} } }],
|
|
29
|
+
])('rejects %s', (_, spec) => {
|
|
30
|
+
expect(() => defineAgent(spec)).toThrow(TypeError)
|
|
31
|
+
})
|
|
32
|
+
|
|
33
|
+
test('builds the model from env when no model is given', () => {
|
|
34
|
+
expect(() => defineAgent({ name: 'A', env })).not.toThrow()
|
|
35
|
+
})
|
|
36
|
+
})
|
|
37
|
+
|
|
38
|
+
describe('defineAgent: ask', () => {
|
|
39
|
+
test('invalid input is terminal', async () => {
|
|
40
|
+
const agent = defineAgent({ name: 'A', model: okModel() })
|
|
41
|
+
await expect(agent.service.ask(fakeCtx(), {})).rejects.toThrow(TerminalError)
|
|
42
|
+
await expect(
|
|
43
|
+
agent.service.ask(fakeCtx(), { prompt: 'hi', notAField: true }),
|
|
44
|
+
).rejects.toThrow(TerminalError)
|
|
45
|
+
})
|
|
46
|
+
|
|
47
|
+
test('sends the caller’s prompt, not a value from loadContext', async () => {
|
|
48
|
+
const model = new MockLanguageModelV4({ doGenerate: textResult('Sure.') })
|
|
49
|
+
const agent = defineAgent({
|
|
50
|
+
name: 'A',
|
|
51
|
+
model,
|
|
52
|
+
loadContext: async () => ({ systemPrompt: '', skills: [], prompt: 'HIJACKED' }),
|
|
53
|
+
})
|
|
54
|
+
|
|
55
|
+
expect(await agent.service.ask(fakeCtx(), { prompt: 'REAL PROMPT' })).toEqual({
|
|
56
|
+
text: 'Sure.',
|
|
57
|
+
steps: 1,
|
|
58
|
+
})
|
|
59
|
+
const sent = JSON.stringify(model.doGenerateCalls[0].prompt)
|
|
60
|
+
expect(sent).toContain('REAL PROMPT')
|
|
61
|
+
expect(sent).not.toContain('HIJACKED')
|
|
62
|
+
})
|
|
63
|
+
|
|
64
|
+
test('load-context journals the system prompt and skill index, never skill bodies', async () => {
|
|
65
|
+
const ctx = fakeCtx()
|
|
66
|
+
const model = new MockLanguageModelV4({
|
|
67
|
+
doGenerate: sequence(
|
|
68
|
+
toolCallResult([{ toolName: 'load_skill', input: { name: 'summarise' } }]),
|
|
69
|
+
textResult('ok'),
|
|
70
|
+
),
|
|
71
|
+
})
|
|
72
|
+
const agent = defineAgent({
|
|
73
|
+
name: 'A',
|
|
74
|
+
model,
|
|
75
|
+
system: 'Static.',
|
|
76
|
+
loadContext: async () => ({
|
|
77
|
+
systemPrompt: 'Be terse.',
|
|
78
|
+
skills: [{ name: 'summarise', description: 'Summarise', body: 'SECRET BODY' }],
|
|
79
|
+
}),
|
|
80
|
+
})
|
|
81
|
+
|
|
82
|
+
await agent.service.ask(ctx, { prompt: 'hi' })
|
|
83
|
+
|
|
84
|
+
const loadContext = ctx.steps.find((s) => s.name === 'load-context')
|
|
85
|
+
expect(loadContext.result).toEqual({
|
|
86
|
+
systemPrompt: 'Be terse.',
|
|
87
|
+
skills: [{ name: 'summarise', description: 'Summarise' }],
|
|
88
|
+
})
|
|
89
|
+
expect(JSON.stringify(model.doGenerateCalls[0].prompt)).toContain('Static.\\n\\nBe terse.')
|
|
90
|
+
expect(JSON.stringify(model.doGenerateCalls[1].prompt)).toContain('SECRET BODY')
|
|
91
|
+
expect(JSON.stringify(ctx.steps.map((s) => s.result))).not.toContain('SECRET BODY')
|
|
92
|
+
})
|
|
93
|
+
|
|
94
|
+
test('a custom output schema returns final_answer’s arguments, validated', async () => {
|
|
95
|
+
const output = {
|
|
96
|
+
type: 'object',
|
|
97
|
+
properties: { score: { type: 'number' } },
|
|
98
|
+
required: ['score'],
|
|
99
|
+
additionalProperties: false,
|
|
100
|
+
}
|
|
101
|
+
const answer = (input) =>
|
|
102
|
+
defineAgent({
|
|
103
|
+
name: 'A',
|
|
104
|
+
output,
|
|
105
|
+
model: new MockLanguageModelV4({
|
|
106
|
+
doGenerate: toolCallResult([{ toolName: 'final_answer', input }]),
|
|
107
|
+
}),
|
|
108
|
+
}).service.ask(fakeCtx(), { prompt: 'rate' })
|
|
109
|
+
|
|
110
|
+
expect(await answer({ score: 7 })).toEqual({ score: 7 })
|
|
111
|
+
await expect(answer({ score: 'high' })).rejects.toThrow(TerminalError)
|
|
112
|
+
})
|
|
113
|
+
|
|
114
|
+
test('a custom input schema with a prompt builder', async () => {
|
|
115
|
+
const model = okModel()
|
|
116
|
+
const agent = defineAgent({
|
|
117
|
+
name: 'A',
|
|
118
|
+
model,
|
|
119
|
+
input: { type: 'object', properties: { isin: { type: 'string' } }, required: ['isin'] },
|
|
120
|
+
prompt: ({ isin }) => `Triage ${isin}`,
|
|
121
|
+
})
|
|
122
|
+
await agent.service.ask(fakeCtx(), { isin: 'CH0012' })
|
|
123
|
+
expect(JSON.stringify(model.doGenerateCalls[0].prompt)).toContain('Triage CH0012')
|
|
124
|
+
})
|
|
125
|
+
|
|
126
|
+
test('createAgentService keeps the old AgentService shape', async () => {
|
|
127
|
+
const agent = createAgentService({ env, model: okModel() })
|
|
128
|
+
expect(agent.name).toBe('AgentService')
|
|
129
|
+
expect(await agent.service.ask(fakeCtx(), { prompt: 'hi' })).toEqual({
|
|
130
|
+
text: 'ok',
|
|
131
|
+
steps: 1,
|
|
132
|
+
})
|
|
133
|
+
})
|
|
134
|
+
})
|
|
135
|
+
|
|
136
|
+
describe('defineAgent: session memory', () => {
|
|
137
|
+
test('without a store: later turns see earlier ones; reset forgets them', async () => {
|
|
138
|
+
const model = new MockLanguageModelV4({
|
|
139
|
+
doGenerate: sequence(textResult('first answer'), textResult('second answer')),
|
|
140
|
+
})
|
|
141
|
+
const agent = defineAgent({ name: 'Chat', model, memory: 'session' })
|
|
142
|
+
const ctx = fakeCtx({ key: 'thread-1' })
|
|
143
|
+
|
|
144
|
+
await agent.object.ask(ctx, { prompt: 'first question' })
|
|
145
|
+
await agent.object.ask(ctx, { prompt: 'second question' })
|
|
146
|
+
|
|
147
|
+
const second = JSON.stringify(model.doGenerateCalls[1].prompt)
|
|
148
|
+
expect(second).toContain('first question')
|
|
149
|
+
expect(second).toContain('first answer')
|
|
150
|
+
expect(ctx.state.history).toHaveLength(2)
|
|
151
|
+
|
|
152
|
+
await agent.object.reset(ctx)
|
|
153
|
+
expect(ctx.state.history).toBeUndefined()
|
|
154
|
+
})
|
|
155
|
+
|
|
156
|
+
test('without a store: keeps only the last maxHistoryTurns turns', async () => {
|
|
157
|
+
const agent = defineAgent({
|
|
158
|
+
name: 'Chat',
|
|
159
|
+
model: okModel(),
|
|
160
|
+
memory: 'session',
|
|
161
|
+
limits: { maxHistoryTurns: 2 },
|
|
162
|
+
})
|
|
163
|
+
const ctx = fakeCtx({ key: 'thread-1' })
|
|
164
|
+
for (const prompt of ['one', 'two', 'three']) await agent.object.ask(ctx, { prompt })
|
|
165
|
+
|
|
166
|
+
expect(ctx.state.history).toHaveLength(2)
|
|
167
|
+
expect(JSON.stringify(ctx.state.history)).not.toContain('"one"')
|
|
168
|
+
})
|
|
169
|
+
|
|
170
|
+
test('with a store: state and journal hold only the ref', async () => {
|
|
171
|
+
const snapshots = new Map()
|
|
172
|
+
const historyStore = {
|
|
173
|
+
load: async (ref) => snapshots.get(ref),
|
|
174
|
+
save: async (sessionKey, messages) => {
|
|
175
|
+
const ref = `${sessionKey}@${snapshots.size + 1}`
|
|
176
|
+
snapshots.set(ref, messages)
|
|
177
|
+
return ref
|
|
178
|
+
},
|
|
179
|
+
}
|
|
180
|
+
const model = new MockLanguageModelV4({
|
|
181
|
+
doGenerate: sequence(textResult('first answer'), textResult('second answer')),
|
|
182
|
+
})
|
|
183
|
+
const agent = defineAgent({ name: 'Chat', model, memory: 'session', historyStore })
|
|
184
|
+
const ctx = fakeCtx({ key: 'thread-1' })
|
|
185
|
+
|
|
186
|
+
await agent.object.ask(ctx, { prompt: 'first question' })
|
|
187
|
+
await agent.object.ask(ctx, { prompt: 'second question' })
|
|
188
|
+
|
|
189
|
+
expect(ctx.state).toEqual({ historyRef: 'thread-1@2' })
|
|
190
|
+
expect(snapshots.get('thread-1@2')).toHaveLength(4)
|
|
191
|
+
expect(JSON.stringify(model.doGenerateCalls[1].prompt)).toContain('first answer')
|
|
192
|
+
const journaled = JSON.stringify(ctx.steps.map((s) => s.result))
|
|
193
|
+
expect(journaled).not.toContain('first question')
|
|
194
|
+
})
|
|
195
|
+
})
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
import { describe, test, expect } from 'bun:test'
|
|
2
|
+
import { TerminalError, RetryableError } from '@restatedev/restate-sdk'
|
|
3
|
+
import { APICallError } from '@ai-sdk/provider'
|
|
4
|
+
import { toRestateLLMError, DEFAULT_LLM_RETRY } from '../../lib/agent/Errors.js'
|
|
5
|
+
|
|
6
|
+
function apiError({ statusCode, isRetryable, responseHeaders } = {}) {
|
|
7
|
+
return new APICallError({
|
|
8
|
+
message: `request failed with status ${statusCode}`,
|
|
9
|
+
url: 'https://llm.example/v1/chat/completions',
|
|
10
|
+
requestBodyValues: {},
|
|
11
|
+
statusCode,
|
|
12
|
+
isRetryable,
|
|
13
|
+
responseHeaders,
|
|
14
|
+
})
|
|
15
|
+
}
|
|
16
|
+
|
|
17
|
+
describe('toRestateLLMError', () => {
|
|
18
|
+
test('non-APICallError errors pass through unchanged (Restate treats them as retryable by default)', () => {
|
|
19
|
+
const original = new TypeError('fetch failed')
|
|
20
|
+
expect(toRestateLLMError(original)).toBe(original)
|
|
21
|
+
})
|
|
22
|
+
|
|
23
|
+
test('a non-retryable APICallError (e.g. 401) becomes a TerminalError carrying the status code', () => {
|
|
24
|
+
const error = apiError({ statusCode: 401, isRetryable: false })
|
|
25
|
+
const mapped = toRestateLLMError(error)
|
|
26
|
+
|
|
27
|
+
expect(mapped).toBeInstanceOf(TerminalError)
|
|
28
|
+
expect(mapped.message).toBe(error.message)
|
|
29
|
+
expect(mapped.code).toBe(401)
|
|
30
|
+
})
|
|
31
|
+
|
|
32
|
+
test('a retryable APICallError with no Retry-After header passes through unchanged', () => {
|
|
33
|
+
const error = apiError({ statusCode: 500, isRetryable: true })
|
|
34
|
+
expect(toRestateLLMError(error)).toBe(error)
|
|
35
|
+
})
|
|
36
|
+
|
|
37
|
+
test('a retryable APICallError with a Retry-After header becomes a RetryableError honoring it', () => {
|
|
38
|
+
const error = apiError({
|
|
39
|
+
statusCode: 429,
|
|
40
|
+
isRetryable: true,
|
|
41
|
+
responseHeaders: { 'retry-after': '20' },
|
|
42
|
+
})
|
|
43
|
+
const mapped = toRestateLLMError(error)
|
|
44
|
+
|
|
45
|
+
expect(mapped).toBeInstanceOf(RetryableError)
|
|
46
|
+
expect(mapped.retryAfter).toEqual({ seconds: 20 })
|
|
47
|
+
})
|
|
48
|
+
|
|
49
|
+
test('a retryable APICallError with a non-numeric Retry-After header passes through unchanged', () => {
|
|
50
|
+
// e.g. an HTTP-date form of Retry-After, which this module doesn't attempt to parse —
|
|
51
|
+
// falling back to ctx.run's own RunOptions backoff is safer than mis-parsing it.
|
|
52
|
+
const error = apiError({
|
|
53
|
+
statusCode: 429,
|
|
54
|
+
isRetryable: true,
|
|
55
|
+
responseHeaders: { 'retry-after': 'Wed, 21 Oct 2026 07:28:00 GMT' },
|
|
56
|
+
})
|
|
57
|
+
expect(toRestateLLMError(error)).toBe(error)
|
|
58
|
+
})
|
|
59
|
+
})
|
|
60
|
+
|
|
61
|
+
describe('DEFAULT_LLM_RETRY', () => {
|
|
62
|
+
test('bounds attempts instead of retrying forever', () => {
|
|
63
|
+
expect(DEFAULT_LLM_RETRY.maxRetryAttempts).toBeGreaterThan(0)
|
|
64
|
+
expect(Number.isFinite(DEFAULT_LLM_RETRY.maxRetryAttempts)).toBe(true)
|
|
65
|
+
})
|
|
66
|
+
})
|