@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,173 @@
|
|
|
1
|
+
import { describe, test, expect } from 'bun:test'
|
|
2
|
+
import { MockLanguageModelV4 } from 'ai/test'
|
|
3
|
+
import { TerminalError, RetryableError } from '@restatedev/restate-sdk'
|
|
4
|
+
import { APICallError } from '@ai-sdk/provider'
|
|
5
|
+
import { runAgentLoop } from '../../lib/agent/Loop.js'
|
|
6
|
+
import { DEFAULT_LLM_RETRY } from '../../lib/agent/Errors.js'
|
|
7
|
+
import { fakeCtx, textResult, toolCallResult, sequence } from './fakes.js'
|
|
8
|
+
|
|
9
|
+
const echoTool = {
|
|
10
|
+
description: 'Echo the input',
|
|
11
|
+
inputSchema: { type: 'object', properties: { v: { type: 'string' } } },
|
|
12
|
+
execute: async ({ v }) => ({ echoed: v }),
|
|
13
|
+
}
|
|
14
|
+
|
|
15
|
+
const apiError = (statusCode, isRetryable, responseHeaders) =>
|
|
16
|
+
new APICallError({
|
|
17
|
+
message: `status ${statusCode}`,
|
|
18
|
+
url: 'https://llm.example/v1/chat/completions',
|
|
19
|
+
requestBodyValues: {},
|
|
20
|
+
statusCode,
|
|
21
|
+
isRetryable,
|
|
22
|
+
responseHeaders,
|
|
23
|
+
})
|
|
24
|
+
|
|
25
|
+
describe('runAgentLoop', () => {
|
|
26
|
+
test('returns the text answer and records each LLM call as its own step', async () => {
|
|
27
|
+
const ctx = fakeCtx()
|
|
28
|
+
const model = new MockLanguageModelV4({ doGenerate: textResult('hi') })
|
|
29
|
+
|
|
30
|
+
const result = await runAgentLoop(ctx, { model, system: '', prompt: 'hello' })
|
|
31
|
+
|
|
32
|
+
expect(result).toMatchObject({ text: 'hi', steps: 1 })
|
|
33
|
+
expect(result.messages[0]).toEqual({ role: 'user', content: 'hello' })
|
|
34
|
+
expect(ctx.steps.map((s) => s.name)).toEqual(['llm-step-0'])
|
|
35
|
+
})
|
|
36
|
+
|
|
37
|
+
test('applies DEFAULT_LLM_RETRY by default, and a caller override', async () => {
|
|
38
|
+
const model = new MockLanguageModelV4({ doGenerate: textResult('hi') })
|
|
39
|
+
const ctx = fakeCtx()
|
|
40
|
+
await runAgentLoop(ctx, { model, system: '', prompt: 'hello' })
|
|
41
|
+
expect(ctx.steps[0].options).toBe(DEFAULT_LLM_RETRY)
|
|
42
|
+
|
|
43
|
+
const retry = { maxRetryAttempts: 1 }
|
|
44
|
+
const ctx2 = fakeCtx()
|
|
45
|
+
await runAgentLoop(ctx2, { model, system: '', prompt: 'hello', retry })
|
|
46
|
+
expect(ctx2.steps[0].options).toBe(retry)
|
|
47
|
+
})
|
|
48
|
+
|
|
49
|
+
test('a non-retryable model failure is terminal', async () => {
|
|
50
|
+
const model = new MockLanguageModelV4({
|
|
51
|
+
doGenerate: () => {
|
|
52
|
+
throw apiError(401, false)
|
|
53
|
+
},
|
|
54
|
+
})
|
|
55
|
+
await expect(
|
|
56
|
+
runAgentLoop(fakeCtx(), { model, system: '', prompt: 'hello' }),
|
|
57
|
+
).rejects.toThrow(TerminalError)
|
|
58
|
+
})
|
|
59
|
+
|
|
60
|
+
test('a rate-limited model failure with Retry-After becomes a RetryableError', async () => {
|
|
61
|
+
const model = new MockLanguageModelV4({
|
|
62
|
+
doGenerate: () => {
|
|
63
|
+
throw apiError(429, true, { 'retry-after': '5' })
|
|
64
|
+
},
|
|
65
|
+
})
|
|
66
|
+
let caught
|
|
67
|
+
try {
|
|
68
|
+
await runAgentLoop(fakeCtx(), { model, system: '', prompt: 'hello', retry: {} })
|
|
69
|
+
} catch (error) {
|
|
70
|
+
caught = error
|
|
71
|
+
}
|
|
72
|
+
expect(caught).toBeInstanceOf(RetryableError)
|
|
73
|
+
expect(caught.retryAfter).toEqual({ seconds: 5 })
|
|
74
|
+
})
|
|
75
|
+
|
|
76
|
+
test('runs a tool, feeds its result back, and continues', async () => {
|
|
77
|
+
const ctx = fakeCtx()
|
|
78
|
+
const model = new MockLanguageModelV4({
|
|
79
|
+
doGenerate: sequence(
|
|
80
|
+
toolCallResult([{ toolName: 'echo', input: { v: 'x' } }]),
|
|
81
|
+
textResult('done'),
|
|
82
|
+
),
|
|
83
|
+
})
|
|
84
|
+
|
|
85
|
+
const result = await runAgentLoop(ctx, {
|
|
86
|
+
model,
|
|
87
|
+
system: '',
|
|
88
|
+
prompt: 'hello',
|
|
89
|
+
tools: { echo: echoTool },
|
|
90
|
+
})
|
|
91
|
+
|
|
92
|
+
expect(result).toMatchObject({ text: 'done', steps: 2 })
|
|
93
|
+
expect(ctx.steps.map((s) => s.name)).toEqual(['llm-step-0', 'tool:echo-0.0', 'llm-step-1'])
|
|
94
|
+
expect(JSON.stringify(model.doGenerateCalls[1].prompt)).toContain('echoed')
|
|
95
|
+
})
|
|
96
|
+
|
|
97
|
+
test('exceeding maxSteps is terminal', async () => {
|
|
98
|
+
const model = new MockLanguageModelV4({
|
|
99
|
+
doGenerate: toolCallResult([{ toolName: 'echo', input: { v: 'x' } }]),
|
|
100
|
+
})
|
|
101
|
+
await expect(
|
|
102
|
+
runAgentLoop(fakeCtx(), {
|
|
103
|
+
model,
|
|
104
|
+
system: '',
|
|
105
|
+
prompt: 'hello',
|
|
106
|
+
tools: { echo: echoTool },
|
|
107
|
+
maxSteps: 1,
|
|
108
|
+
}),
|
|
109
|
+
).rejects.toThrow(/maxSteps/)
|
|
110
|
+
})
|
|
111
|
+
|
|
112
|
+
test('exceeding maxTokens is terminal with code 429', async () => {
|
|
113
|
+
const model = new MockLanguageModelV4({ doGenerate: textResult('hi') })
|
|
114
|
+
let caught
|
|
115
|
+
try {
|
|
116
|
+
await runAgentLoop(fakeCtx(), { model, system: '', prompt: 'hello', maxTokens: 1 })
|
|
117
|
+
} catch (error) {
|
|
118
|
+
caught = error
|
|
119
|
+
}
|
|
120
|
+
expect(caught).toBeInstanceOf(TerminalError)
|
|
121
|
+
expect(caught.code).toBe(429)
|
|
122
|
+
})
|
|
123
|
+
|
|
124
|
+
test('puts only the skill index in the system prompt', async () => {
|
|
125
|
+
const model = new MockLanguageModelV4({ doGenerate: textResult('ok') })
|
|
126
|
+
await runAgentLoop(fakeCtx(), {
|
|
127
|
+
model,
|
|
128
|
+
system: 'Be terse.',
|
|
129
|
+
prompt: 'hello',
|
|
130
|
+
skills: [{ name: 'summarise', description: 'Summarise things' }],
|
|
131
|
+
})
|
|
132
|
+
const sent = JSON.stringify(model.doGenerateCalls[0].prompt)
|
|
133
|
+
expect(sent).toContain('Be terse.')
|
|
134
|
+
expect(sent).toContain('summarise: Summarise things')
|
|
135
|
+
})
|
|
136
|
+
|
|
137
|
+
test('with an outputSchema, nudges a plain-text answer and returns final_answer’s input', async () => {
|
|
138
|
+
const ctx = fakeCtx()
|
|
139
|
+
const model = new MockLanguageModelV4({
|
|
140
|
+
doGenerate: sequence(
|
|
141
|
+
textResult('here you go'),
|
|
142
|
+
toolCallResult([{ toolName: 'final_answer', input: { score: 7 } }]),
|
|
143
|
+
),
|
|
144
|
+
})
|
|
145
|
+
|
|
146
|
+
const result = await runAgentLoop(ctx, {
|
|
147
|
+
model,
|
|
148
|
+
system: '',
|
|
149
|
+
prompt: 'rate it',
|
|
150
|
+
outputSchema: { type: 'object', properties: { score: { type: 'number' } } },
|
|
151
|
+
})
|
|
152
|
+
|
|
153
|
+
expect(result).toMatchObject({ output: { score: 7 }, steps: 2 })
|
|
154
|
+
expect(JSON.stringify(model.doGenerateCalls[1].prompt)).toContain('final_answer')
|
|
155
|
+
})
|
|
156
|
+
|
|
157
|
+
test('prepends loadHistory’s messages to every model call, without journaling them', async () => {
|
|
158
|
+
const ctx = fakeCtx()
|
|
159
|
+
const model = new MockLanguageModelV4({ doGenerate: textResult('ok') })
|
|
160
|
+
const history = [{ role: 'user', content: 'EARLIER TURN' }]
|
|
161
|
+
|
|
162
|
+
const result = await runAgentLoop(ctx, {
|
|
163
|
+
model,
|
|
164
|
+
system: '',
|
|
165
|
+
prompt: 'hello',
|
|
166
|
+
loadHistory: async () => history,
|
|
167
|
+
})
|
|
168
|
+
|
|
169
|
+
expect(JSON.stringify(model.doGenerateCalls[0].prompt)).toContain('EARLIER TURN')
|
|
170
|
+
expect(JSON.stringify(ctx.steps[0].result)).not.toContain('EARLIER TURN')
|
|
171
|
+
expect(JSON.stringify(result.messages)).not.toContain('EARLIER TURN')
|
|
172
|
+
})
|
|
173
|
+
})
|
|
@@ -0,0 +1,147 @@
|
|
|
1
|
+
import { describe, test, expect } from 'bun:test'
|
|
2
|
+
import { MockLanguageModelV4 } from 'ai/test'
|
|
3
|
+
import { TerminalError } from '@restatedev/restate-sdk'
|
|
4
|
+
import {
|
|
5
|
+
cachedSkillFetcher,
|
|
6
|
+
expandSkillRefs,
|
|
7
|
+
hashSkill,
|
|
8
|
+
skillRefOutput,
|
|
9
|
+
} from '../../lib/agent/Skills.js'
|
|
10
|
+
import { runAgentLoop } from '../../lib/agent/Loop.js'
|
|
11
|
+
import { fakeCtx, sequence, textResult, toolCallResult } from './fakes.js'
|
|
12
|
+
|
|
13
|
+
const skillMessage = (body) => ({
|
|
14
|
+
role: 'tool',
|
|
15
|
+
content: [
|
|
16
|
+
{
|
|
17
|
+
type: 'tool-result',
|
|
18
|
+
toolCallId: 'call_0',
|
|
19
|
+
toolName: 'load_skill',
|
|
20
|
+
output: skillRefOutput({ name: 'summarise', sha256: hashSkill(body) }),
|
|
21
|
+
},
|
|
22
|
+
],
|
|
23
|
+
})
|
|
24
|
+
const fetchFrom = (skills) => async (name) => skills[name]
|
|
25
|
+
|
|
26
|
+
describe('expandSkillRefs', () => {
|
|
27
|
+
test('replaces a reference with the body when the hash matches', async () => {
|
|
28
|
+
const [expanded] = await expandSkillRefs([skillMessage('v1')], {
|
|
29
|
+
fetchSkill: fetchFrom({ summarise: 'v1' }),
|
|
30
|
+
strict: true,
|
|
31
|
+
})
|
|
32
|
+
expect(expanded.content[0].output).toEqual({ type: 'text', value: 'v1' })
|
|
33
|
+
})
|
|
34
|
+
|
|
35
|
+
test('leaves other messages and tool results untouched', async () => {
|
|
36
|
+
const messages = [
|
|
37
|
+
{ role: 'user', content: 'hi' },
|
|
38
|
+
{
|
|
39
|
+
role: 'tool',
|
|
40
|
+
content: [
|
|
41
|
+
{
|
|
42
|
+
type: 'tool-result',
|
|
43
|
+
toolCallId: 'c',
|
|
44
|
+
toolName: 'echo',
|
|
45
|
+
output: { type: 'json', value: 1 },
|
|
46
|
+
},
|
|
47
|
+
],
|
|
48
|
+
},
|
|
49
|
+
]
|
|
50
|
+
expect(
|
|
51
|
+
await expandSkillRefs(messages, { fetchSkill: fetchFrom({}), strict: true }),
|
|
52
|
+
).toEqual(messages)
|
|
53
|
+
})
|
|
54
|
+
|
|
55
|
+
test('strict: a changed or missing skill is terminal', async () => {
|
|
56
|
+
for (const skills of [{ summarise: 'v2' }, {}]) {
|
|
57
|
+
await expect(
|
|
58
|
+
expandSkillRefs([skillMessage('v1')], {
|
|
59
|
+
fetchSkill: fetchFrom(skills),
|
|
60
|
+
strict: true,
|
|
61
|
+
}),
|
|
62
|
+
).rejects.toBeInstanceOf(TerminalError)
|
|
63
|
+
}
|
|
64
|
+
})
|
|
65
|
+
|
|
66
|
+
test('lenient: a changed skill uses the current body; a missing one becomes a notice', async () => {
|
|
67
|
+
const [changed] = await expandSkillRefs([skillMessage('v1')], {
|
|
68
|
+
fetchSkill: fetchFrom({ summarise: 'v2' }),
|
|
69
|
+
strict: false,
|
|
70
|
+
})
|
|
71
|
+
expect(changed.content[0].output).toEqual({ type: 'text', value: 'v2' })
|
|
72
|
+
|
|
73
|
+
const [missing] = await expandSkillRefs([skillMessage('v1')], {
|
|
74
|
+
fetchSkill: undefined,
|
|
75
|
+
strict: false,
|
|
76
|
+
})
|
|
77
|
+
expect(missing.content[0].output).toEqual({
|
|
78
|
+
type: 'error-text',
|
|
79
|
+
value: 'Skill "summarise" is no longer available.',
|
|
80
|
+
})
|
|
81
|
+
})
|
|
82
|
+
})
|
|
83
|
+
|
|
84
|
+
describe('cachedSkillFetcher', () => {
|
|
85
|
+
test('fetches each skill once per execution', async () => {
|
|
86
|
+
let fetches = 0
|
|
87
|
+
const fetchSkill = cachedSkillFetcher(async () => {
|
|
88
|
+
fetches++
|
|
89
|
+
return 'body'
|
|
90
|
+
})
|
|
91
|
+
await fetchSkill('a')
|
|
92
|
+
await fetchSkill('a')
|
|
93
|
+
await fetchSkill('b')
|
|
94
|
+
expect(fetches).toBe(2)
|
|
95
|
+
})
|
|
96
|
+
|
|
97
|
+
test('is undefined without a loader', () => {
|
|
98
|
+
expect(cachedSkillFetcher(undefined)).toBeUndefined()
|
|
99
|
+
})
|
|
100
|
+
})
|
|
101
|
+
|
|
102
|
+
describe('runAgentLoop with skills', () => {
|
|
103
|
+
const skills = [{ name: 'summarise', description: 'Summarise' }]
|
|
104
|
+
|
|
105
|
+
test('the model sees the body; the journal holds only the reference', async () => {
|
|
106
|
+
const ctx = fakeCtx()
|
|
107
|
+
const model = new MockLanguageModelV4({
|
|
108
|
+
doGenerate: sequence(
|
|
109
|
+
toolCallResult([{ toolName: 'load_skill', input: { name: 'summarise' } }]),
|
|
110
|
+
textResult('done'),
|
|
111
|
+
),
|
|
112
|
+
})
|
|
113
|
+
|
|
114
|
+
const result = await runAgentLoop(ctx, {
|
|
115
|
+
model,
|
|
116
|
+
system: '',
|
|
117
|
+
prompt: 'hi',
|
|
118
|
+
skills,
|
|
119
|
+
loadSkill: async () => 'FULL SKILL BODY',
|
|
120
|
+
})
|
|
121
|
+
|
|
122
|
+
expect(JSON.stringify(model.doGenerateCalls[1].prompt)).toContain('FULL SKILL BODY')
|
|
123
|
+
expect(JSON.stringify(ctx.steps.map((s) => s.result))).not.toContain('FULL SKILL BODY')
|
|
124
|
+
expect(JSON.stringify(result.messages)).not.toContain('FULL SKILL BODY')
|
|
125
|
+
})
|
|
126
|
+
|
|
127
|
+
test('a skill that changes mid-invocation fails the step instead of changing the instructions', async () => {
|
|
128
|
+
let version = 0
|
|
129
|
+
const model = new MockLanguageModelV4({
|
|
130
|
+
doGenerate: sequence(
|
|
131
|
+
toolCallResult([{ toolName: 'load_skill', input: { name: 'summarise' } }]),
|
|
132
|
+
textResult('done'),
|
|
133
|
+
),
|
|
134
|
+
})
|
|
135
|
+
|
|
136
|
+
await expect(
|
|
137
|
+
runAgentLoop(fakeCtx(), {
|
|
138
|
+
model,
|
|
139
|
+
system: '',
|
|
140
|
+
prompt: 'hi',
|
|
141
|
+
skills,
|
|
142
|
+
// The journaled hash is taken from v1; the next LLM step fetches v2.
|
|
143
|
+
loadSkill: async () => `v${++version}`,
|
|
144
|
+
}),
|
|
145
|
+
).rejects.toThrow(/changed or disappeared/)
|
|
146
|
+
})
|
|
147
|
+
})
|
|
@@ -0,0 +1,277 @@
|
|
|
1
|
+
import { describe, test, expect } from 'bun:test'
|
|
2
|
+
import { CancelledError, TerminalError, TimeoutError, serde } from '@restatedev/restate-sdk'
|
|
3
|
+
import { assertToolSpecs, executeTool } from '../../lib/agent/Tools.js'
|
|
4
|
+
import { DEFAULT_TOOL_RETRY } from '../../lib/agent/Errors.js'
|
|
5
|
+
import { hashSkill } from '../../lib/agent/Skills.js'
|
|
6
|
+
import { fakeCtx } from './fakes.js'
|
|
7
|
+
|
|
8
|
+
const schema = { type: 'object' }
|
|
9
|
+
const call = (toolName, input = {}) => ({ toolCallId: 'call_0', toolName, input })
|
|
10
|
+
const at = { step: 2, index: 1 }
|
|
11
|
+
|
|
12
|
+
describe('assertToolSpecs', () => {
|
|
13
|
+
test('accepts block and execute tools', () => {
|
|
14
|
+
expect(() =>
|
|
15
|
+
assertToolSpecs({
|
|
16
|
+
a: { description: 'a', inputSchema: schema, block: 'Svc', handler: 'do' },
|
|
17
|
+
b: { description: 'b', inputSchema: schema, execute: async () => 1 },
|
|
18
|
+
}),
|
|
19
|
+
).not.toThrow()
|
|
20
|
+
})
|
|
21
|
+
|
|
22
|
+
test.each([
|
|
23
|
+
[
|
|
24
|
+
'a reserved name',
|
|
25
|
+
{ final_answer: { description: 'x', inputSchema: schema, execute() {} } },
|
|
26
|
+
],
|
|
27
|
+
['an invalid name', { 'a b': { description: 'x', inputSchema: schema, execute() {} } }],
|
|
28
|
+
['no inputSchema', { a: { description: 'x', execute() {} } }],
|
|
29
|
+
[
|
|
30
|
+
'both block and execute',
|
|
31
|
+
{
|
|
32
|
+
a: {
|
|
33
|
+
description: 'x',
|
|
34
|
+
inputSchema: schema,
|
|
35
|
+
block: 'S',
|
|
36
|
+
handler: 'h',
|
|
37
|
+
execute() {},
|
|
38
|
+
},
|
|
39
|
+
},
|
|
40
|
+
],
|
|
41
|
+
['a block without handler', { a: { description: 'x', inputSchema: schema, block: 'S' } }],
|
|
42
|
+
[
|
|
43
|
+
'approval without notify',
|
|
44
|
+
{
|
|
45
|
+
a: {
|
|
46
|
+
description: 'x',
|
|
47
|
+
inputSchema: schema,
|
|
48
|
+
execute() {},
|
|
49
|
+
approval: { timeout: 10 },
|
|
50
|
+
},
|
|
51
|
+
},
|
|
52
|
+
],
|
|
53
|
+
])('rejects %s', (_, tools) => {
|
|
54
|
+
expect(() => assertToolSpecs(tools)).toThrow(TypeError)
|
|
55
|
+
})
|
|
56
|
+
})
|
|
57
|
+
|
|
58
|
+
describe('executeTool', () => {
|
|
59
|
+
test('execute tool: own bounded step, with an idempotency key', async () => {
|
|
60
|
+
const ctx = fakeCtx()
|
|
61
|
+
let meta
|
|
62
|
+
const tools = {
|
|
63
|
+
write: {
|
|
64
|
+
description: 'w',
|
|
65
|
+
inputSchema: schema,
|
|
66
|
+
execute: async (args, m) => {
|
|
67
|
+
meta = m
|
|
68
|
+
return { wrote: args.v }
|
|
69
|
+
},
|
|
70
|
+
},
|
|
71
|
+
}
|
|
72
|
+
|
|
73
|
+
const output = await executeTool(ctx, call('write', { v: 1 }), { tools, ...at })
|
|
74
|
+
|
|
75
|
+
expect(output).toEqual({ type: 'json', value: { wrote: 1 } })
|
|
76
|
+
expect(meta).toEqual({ idempotencyKey: 'inv-1:2.1' })
|
|
77
|
+
expect(ctx.steps[0]).toMatchObject({ name: 'tool:write-2.1', options: DEFAULT_TOOL_RETRY })
|
|
78
|
+
})
|
|
79
|
+
|
|
80
|
+
test('execute tool: exhausted retries reach the model as error-text', async () => {
|
|
81
|
+
const tools = {
|
|
82
|
+
flaky: {
|
|
83
|
+
description: 'f',
|
|
84
|
+
inputSchema: schema,
|
|
85
|
+
execute: async () => {
|
|
86
|
+
throw new Error('connection reset')
|
|
87
|
+
},
|
|
88
|
+
},
|
|
89
|
+
}
|
|
90
|
+
const output = await executeTool(fakeCtx(), call('flaky'), { tools, ...at })
|
|
91
|
+
expect(output).toEqual({ type: 'error-text', value: 'connection reset' })
|
|
92
|
+
})
|
|
93
|
+
|
|
94
|
+
test('execute tool without a retry limit: a transient error propagates, so Restate retries', async () => {
|
|
95
|
+
const tools = {
|
|
96
|
+
flaky: {
|
|
97
|
+
description: 'f',
|
|
98
|
+
inputSchema: schema,
|
|
99
|
+
retry: {},
|
|
100
|
+
execute: async () => {
|
|
101
|
+
throw new Error('connection reset')
|
|
102
|
+
},
|
|
103
|
+
},
|
|
104
|
+
}
|
|
105
|
+
await expect(executeTool(fakeCtx(), call('flaky'), { tools, ...at })).rejects.toThrow(
|
|
106
|
+
'connection reset',
|
|
107
|
+
)
|
|
108
|
+
})
|
|
109
|
+
|
|
110
|
+
test('block tool: a Restate call to the target handler', async () => {
|
|
111
|
+
const ctx = fakeCtx({ onCall: () => ({ id: 'ISS-1' }) })
|
|
112
|
+
const tools = {
|
|
113
|
+
lookup: {
|
|
114
|
+
description: 'l',
|
|
115
|
+
inputSchema: schema,
|
|
116
|
+
block: { name: 'IssuerService' },
|
|
117
|
+
handler: 'get',
|
|
118
|
+
key: (args) => args.isin,
|
|
119
|
+
},
|
|
120
|
+
}
|
|
121
|
+
|
|
122
|
+
const output = await executeTool(ctx, call('lookup', { isin: 'CH01' }), { tools, ...at })
|
|
123
|
+
|
|
124
|
+
expect(output).toEqual({ type: 'json', value: { id: 'ISS-1' } })
|
|
125
|
+
expect(ctx.calls[0]).toEqual({
|
|
126
|
+
service: 'IssuerService',
|
|
127
|
+
method: 'get',
|
|
128
|
+
key: 'CH01',
|
|
129
|
+
parameter: { isin: 'CH01' },
|
|
130
|
+
inputSerde: serde.json,
|
|
131
|
+
outputSerde: serde.json,
|
|
132
|
+
})
|
|
133
|
+
expect(ctx.steps).toEqual([])
|
|
134
|
+
})
|
|
135
|
+
|
|
136
|
+
test('block tool: a terminal error from the callee reaches the model as error-text', async () => {
|
|
137
|
+
const ctx = fakeCtx({
|
|
138
|
+
onCall: () => {
|
|
139
|
+
throw new TerminalError('issuer not found')
|
|
140
|
+
},
|
|
141
|
+
})
|
|
142
|
+
const tools = {
|
|
143
|
+
lookup: {
|
|
144
|
+
description: 'l',
|
|
145
|
+
inputSchema: schema,
|
|
146
|
+
block: 'IssuerService',
|
|
147
|
+
handler: 'get',
|
|
148
|
+
},
|
|
149
|
+
}
|
|
150
|
+
expect(await executeTool(ctx, call('lookup'), { tools, ...at })).toEqual({
|
|
151
|
+
type: 'error-text',
|
|
152
|
+
value: 'issuer not found',
|
|
153
|
+
})
|
|
154
|
+
})
|
|
155
|
+
|
|
156
|
+
test('cancellation always propagates', async () => {
|
|
157
|
+
const ctx = fakeCtx({
|
|
158
|
+
onCall: () => {
|
|
159
|
+
throw new CancelledError()
|
|
160
|
+
},
|
|
161
|
+
})
|
|
162
|
+
const tools = {
|
|
163
|
+
lookup: {
|
|
164
|
+
description: 'l',
|
|
165
|
+
inputSchema: schema,
|
|
166
|
+
block: 'IssuerService',
|
|
167
|
+
handler: 'get',
|
|
168
|
+
},
|
|
169
|
+
}
|
|
170
|
+
await expect(executeTool(ctx, call('lookup'), { tools, ...at })).rejects.toBeInstanceOf(
|
|
171
|
+
CancelledError,
|
|
172
|
+
)
|
|
173
|
+
})
|
|
174
|
+
|
|
175
|
+
test('unknown tool reaches the model as error-text', async () => {
|
|
176
|
+
expect(await executeTool(fakeCtx(), call('nope'), { tools: {}, ...at })).toEqual({
|
|
177
|
+
type: 'error-text',
|
|
178
|
+
value: 'Unknown tool "nope".',
|
|
179
|
+
})
|
|
180
|
+
})
|
|
181
|
+
|
|
182
|
+
test('load_skill journals only a reference and hash, never the body', async () => {
|
|
183
|
+
const ctx = fakeCtx()
|
|
184
|
+
const loadSkill = async (name) => (name === 'summarise' ? 'FULL BODY' : undefined)
|
|
185
|
+
const skillRef = { name: 'summarise', sha256: hashSkill('FULL BODY') }
|
|
186
|
+
|
|
187
|
+
expect(
|
|
188
|
+
await executeTool(ctx, call('load_skill', { name: 'summarise' }), {
|
|
189
|
+
tools: {},
|
|
190
|
+
loadSkill,
|
|
191
|
+
...at,
|
|
192
|
+
}),
|
|
193
|
+
).toEqual({ type: 'json', value: { skillRef } })
|
|
194
|
+
expect(ctx.steps[0]).toMatchObject({ name: 'skill:summarise-2.1', result: skillRef })
|
|
195
|
+
|
|
196
|
+
expect(
|
|
197
|
+
await executeTool(ctx, call('load_skill', { name: 'other' }), {
|
|
198
|
+
tools: {},
|
|
199
|
+
loadSkill,
|
|
200
|
+
...at,
|
|
201
|
+
}),
|
|
202
|
+
).toEqual({ type: 'error-text', value: 'No skill named "other" is available.' })
|
|
203
|
+
})
|
|
204
|
+
|
|
205
|
+
describe('approval', () => {
|
|
206
|
+
const approvalTools = (executed) => ({
|
|
207
|
+
notifyDesk: {
|
|
208
|
+
description: 'n',
|
|
209
|
+
inputSchema: schema,
|
|
210
|
+
execute: async () => {
|
|
211
|
+
executed.push(true)
|
|
212
|
+
return 'sent'
|
|
213
|
+
},
|
|
214
|
+
approval: { timeout: 1000, notify: { block: 'ApprovalInbox', handler: 'request' } },
|
|
215
|
+
},
|
|
216
|
+
})
|
|
217
|
+
|
|
218
|
+
test('notifies the approver with the awakeable ID, then runs the tool once approved', async () => {
|
|
219
|
+
const executed = []
|
|
220
|
+
let waited
|
|
221
|
+
const ctx = fakeCtx({
|
|
222
|
+
onApproval: (timeout) => {
|
|
223
|
+
waited = timeout
|
|
224
|
+
return { approved: true }
|
|
225
|
+
},
|
|
226
|
+
})
|
|
227
|
+
|
|
228
|
+
const output = await executeTool(ctx, call('notifyDesk', { msg: 'hi' }), {
|
|
229
|
+
tools: approvalTools(executed),
|
|
230
|
+
...at,
|
|
231
|
+
})
|
|
232
|
+
|
|
233
|
+
expect(output).toEqual({ type: 'json', value: 'sent' })
|
|
234
|
+
expect(executed).toHaveLength(1)
|
|
235
|
+
expect(waited).toBe(1000)
|
|
236
|
+
expect(ctx.sends[0]).toMatchObject({
|
|
237
|
+
service: 'ApprovalInbox',
|
|
238
|
+
method: 'request',
|
|
239
|
+
parameter: {
|
|
240
|
+
awakeableId: 'awk-1',
|
|
241
|
+
tool: 'notifyDesk',
|
|
242
|
+
input: { msg: 'hi' },
|
|
243
|
+
invocationId: 'inv-1',
|
|
244
|
+
},
|
|
245
|
+
})
|
|
246
|
+
})
|
|
247
|
+
|
|
248
|
+
test('a rejection skips the tool and tells the model why', async () => {
|
|
249
|
+
const executed = []
|
|
250
|
+
const ctx = fakeCtx({ onApproval: () => ({ approved: false, reason: 'not today' }) })
|
|
251
|
+
const output = await executeTool(ctx, call('notifyDesk'), {
|
|
252
|
+
tools: approvalTools(executed),
|
|
253
|
+
...at,
|
|
254
|
+
})
|
|
255
|
+
expect(output).toEqual({
|
|
256
|
+
type: 'error-text',
|
|
257
|
+
value: 'Tool "notifyDesk" was not approved: not today.',
|
|
258
|
+
})
|
|
259
|
+
expect(executed).toHaveLength(0)
|
|
260
|
+
})
|
|
261
|
+
|
|
262
|
+
test('a timeout counts as a rejection', async () => {
|
|
263
|
+
const executed = []
|
|
264
|
+
const ctx = fakeCtx({
|
|
265
|
+
onApproval: () => {
|
|
266
|
+
throw new TimeoutError()
|
|
267
|
+
},
|
|
268
|
+
})
|
|
269
|
+
const output = await executeTool(ctx, call('notifyDesk'), {
|
|
270
|
+
tools: approvalTools(executed),
|
|
271
|
+
...at,
|
|
272
|
+
})
|
|
273
|
+
expect(output.value).toContain('approval timed out')
|
|
274
|
+
expect(executed).toHaveLength(0)
|
|
275
|
+
})
|
|
276
|
+
})
|
|
277
|
+
})
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
import { TerminalError } from '@restatedev/restate-sdk'
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* A stand-in for Restate's `Context` / `ObjectContext`. Per this repo's testing convention,
|
|
5
|
+
* handlers are unit-tested against a mocked `ctx`, never a live Restate server.
|
|
6
|
+
*
|
|
7
|
+
* `run` mimics the parts of the real one the agent depends on: the result goes through a
|
|
8
|
+
* JSON round trip (as it does through the journal), and when a bounded `ctx.run` gives up
|
|
9
|
+
* it throws a `TerminalError` wrapping the original message. One attempt stands in for
|
|
10
|
+
* "all retries exhausted".
|
|
11
|
+
*/
|
|
12
|
+
export function fakeCtx({ key, state = {}, onCall, onApproval } = {}) {
|
|
13
|
+
const steps = []
|
|
14
|
+
const calls = []
|
|
15
|
+
const sends = []
|
|
16
|
+
return {
|
|
17
|
+
key,
|
|
18
|
+
state,
|
|
19
|
+
steps,
|
|
20
|
+
calls,
|
|
21
|
+
sends,
|
|
22
|
+
run: async (name, fn, options) => {
|
|
23
|
+
const step = { name, options }
|
|
24
|
+
steps.push(step)
|
|
25
|
+
let value
|
|
26
|
+
try {
|
|
27
|
+
value = await fn()
|
|
28
|
+
} catch (error) {
|
|
29
|
+
if (error instanceof TerminalError || !options?.maxRetryAttempts) throw error
|
|
30
|
+
throw new TerminalError(error.message)
|
|
31
|
+
}
|
|
32
|
+
step.result = value === undefined ? undefined : JSON.parse(JSON.stringify(value))
|
|
33
|
+
return step.result
|
|
34
|
+
},
|
|
35
|
+
request: () => ({ id: 'inv-1' }),
|
|
36
|
+
genericCall: async (call) => {
|
|
37
|
+
calls.push(call)
|
|
38
|
+
return onCall ? onCall(call) : null
|
|
39
|
+
},
|
|
40
|
+
genericSend: (send) => {
|
|
41
|
+
sends.push(send)
|
|
42
|
+
},
|
|
43
|
+
awakeable: () => ({
|
|
44
|
+
id: 'awk-1',
|
|
45
|
+
promise: { orTimeout: async (timeout) => onApproval?.(timeout) },
|
|
46
|
+
}),
|
|
47
|
+
get: async (name) => state[name],
|
|
48
|
+
set: (name, value) => {
|
|
49
|
+
state[name] = value
|
|
50
|
+
},
|
|
51
|
+
clear: (name) => {
|
|
52
|
+
delete state[name]
|
|
53
|
+
},
|
|
54
|
+
}
|
|
55
|
+
}
|
|
56
|
+
|
|
57
|
+
const usage = { inputTokens: { total: 1 }, outputTokens: { total: 1 } }
|
|
58
|
+
|
|
59
|
+
export function textResult(text) {
|
|
60
|
+
return { content: [{ type: 'text', text }], finishReason: 'stop', usage, warnings: [] }
|
|
61
|
+
}
|
|
62
|
+
|
|
63
|
+
/** @param {Array<{toolName: string, input: object}>} calls */
|
|
64
|
+
export function toolCallResult(calls) {
|
|
65
|
+
return {
|
|
66
|
+
content: calls.map(({ toolName, input }, i) => ({
|
|
67
|
+
type: 'tool-call',
|
|
68
|
+
toolCallId: `call_${i}`,
|
|
69
|
+
toolName,
|
|
70
|
+
input: JSON.stringify(input),
|
|
71
|
+
})),
|
|
72
|
+
finishReason: 'tool-calls',
|
|
73
|
+
usage,
|
|
74
|
+
warnings: [],
|
|
75
|
+
}
|
|
76
|
+
}
|
|
77
|
+
|
|
78
|
+
/** `doGenerate` that returns `results` one after another. */
|
|
79
|
+
export function sequence(...results) {
|
|
80
|
+
let i = 0
|
|
81
|
+
return async () => results[Math.min(i++, results.length - 1)]
|
|
82
|
+
}
|
|
83
|
+
|
|
84
|
+
export const env = { LLM_BASE_URL: 'https://unused.example', MODEL_ID: 'unused' }
|