@statelyai/agent 1.1.6 → 2.0.0-next.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.
Files changed (54) hide show
  1. package/.changeset/light-hats-drive.md +9 -0
  2. package/.changeset/pre.json +10 -0
  3. package/.vscode/launch.json +6 -0
  4. package/CHANGELOG.md +10 -0
  5. package/dist/index.d.mts +262 -165
  6. package/dist/index.d.ts +262 -165
  7. package/dist/index.js +368 -263
  8. package/dist/index.mjs +371 -261
  9. package/examples/chatbot-alt.ts +57 -0
  10. package/examples/chatbot.ts +11 -16
  11. package/examples/cot.ts +25 -22
  12. package/examples/customer-service-sim.ts +107 -0
  13. package/examples/email.ts +14 -14
  14. package/examples/example.ts +5 -5
  15. package/examples/executor.ts +66 -0
  16. package/examples/goal.ts +11 -11
  17. package/examples/helpers/helpers.ts +26 -14
  18. package/examples/joke.ts +78 -75
  19. package/examples/jugs.ts +125 -0
  20. package/examples/multi.ts +4 -4
  21. package/examples/newspaper.ts +98 -104
  22. package/examples/number.ts +5 -4
  23. package/examples/raffle.ts +10 -11
  24. package/examples/river-crossing.ts +140 -0
  25. package/examples/sandbox.ts +1 -1
  26. package/examples/simple.ts +4 -2
  27. package/examples/summary.ts +121 -0
  28. package/examples/support.ts +5 -5
  29. package/examples/ticTacToe.ts +86 -45
  30. package/examples/todo.ts +6 -6
  31. package/examples/tutor.ts +13 -13
  32. package/examples/verify.ts +2 -2
  33. package/examples/weather.ts +5 -8
  34. package/examples/wiki.ts +26 -7
  35. package/examples/word.ts +15 -10
  36. package/package.json +13 -10
  37. package/readme.md +1 -1
  38. package/src/agent-experimental.ts +1 -1
  39. package/src/agent.test.ts +117 -228
  40. package/src/agent.ts +469 -81
  41. package/src/{decision.test.ts → decide.test.ts} +26 -50
  42. package/src/decide.ts +153 -0
  43. package/src/index.ts +1 -1
  44. package/src/middleware.ts +103 -0
  45. package/src/mockModel.ts +47 -0
  46. package/src/planners/shortestPathPlanner.ts +151 -13
  47. package/src/planners/simplePlanner.ts +57 -85
  48. package/src/strategies/chain-of-note.ts +6 -55
  49. package/src/text.ts +51 -144
  50. package/src/types.ts +172 -204
  51. package/src/utils.ts +37 -4
  52. package/src/adapters/vercel.ts +0 -7
  53. package/src/decision.ts +0 -84
  54. package/src/memory.ts +0 -25
@@ -1,4 +1,4 @@
1
- import { type CoreTool, tool } from 'ai';
1
+ import { CoreMessage, type CoreTool, generateText, tool } from 'ai';
2
2
  import {
3
3
  AgentPlan,
4
4
  AgentPlanInput,
@@ -7,22 +7,11 @@ import {
7
7
  TransitionData,
8
8
  AnyAgent,
9
9
  } from '../types';
10
- import { getAllTransitions } from '../utils';
11
- import { AnyStateMachine } from 'xstate';
10
+ import { getAllTransitions, randomId } from '../utils';
11
+ import { AnyStateMachine, getNextSnapshot } from 'xstate';
12
12
  import { defaultTextTemplate } from '../templates/defaultText';
13
13
  import { getMessages } from '../text';
14
-
15
- function getTransitions(
16
- state: ObservedState,
17
- machine: AnyStateMachine
18
- ): TransitionData[] {
19
- if (!machine) {
20
- return [];
21
- }
22
-
23
- const resolvedState = machine.resolveState(state);
24
- return getAllTransitions(resolvedState);
25
- }
14
+ import { getToolMap } from '../decide';
26
15
 
27
16
  const simplePlannerPromptTemplate: PromptTemplate<any> = (data) => {
28
17
  return `
@@ -36,64 +25,9 @@ export async function simplePlanner<T extends AnyAgent>(
36
25
  agent: T,
37
26
  input: AgentPlanInput<any>
38
27
  ): Promise<AgentPlan<any> | undefined> {
39
- // Get all of the possible next transitions
40
- const transitions: TransitionData[] = input.machine
41
- ? getTransitions(input.state, input.machine)
42
- : Object.entries(input.events).map(([eventType, { description }]) => ({
43
- eventType,
44
- description,
45
- }));
46
-
47
- // Only keep the transitions that match the event types that are in the event mapping
48
- // TODO: allow for custom filters
49
- const filter = (eventType: string) =>
50
- Object.keys(input.events).includes(eventType);
51
-
52
- // Mapping of each event type (e.g. "mouse.click")
53
- // to a valid function name (e.g. "mouse_click")
54
- const functionNameMapping: Record<string, string> = {};
55
-
56
- const toolTransitions = transitions
57
- .filter((t) => {
58
- return filter(t.eventType);
59
- })
60
- .map((t) => {
61
- const name = t.eventType.replace(/\./g, '_');
62
- functionNameMapping[name] = t.eventType;
63
-
64
- return {
65
- type: 'function',
66
- eventType: t.eventType,
67
- description: t.description,
68
- name,
69
- } as const;
70
- });
71
-
72
- // Convert the transition data to a tool map that the
73
- // Vercel AI SDK can use
74
- const toolMap: Record<string, CoreTool<any, any>> = {};
75
- for (const toolTransitionData of toolTransitions) {
76
- const toolZodType = input.events?.[toolTransitionData.eventType];
77
-
78
- if (!toolZodType) {
79
- continue;
80
- }
28
+ const toolMap = getToolMap(agent, input);
81
29
 
82
- toolMap[toolTransitionData.name] = tool({
83
- description: toolZodType?.description ?? toolTransitionData.description,
84
- parameters: toolZodType,
85
- execute: async (params: Record<string, any>) => {
86
- const event = {
87
- type: toolTransitionData.eventType,
88
- ...params,
89
- };
90
-
91
- return event;
92
- },
93
- });
94
- }
95
-
96
- if (!Object.keys(toolMap).length) {
30
+ if (!toolMap) {
97
31
  // No valid transitions for the specified tools
98
32
  return undefined;
99
33
  }
@@ -107,12 +41,41 @@ export async function simplePlanner<T extends AnyAgent>(
107
41
 
108
42
  const messages = await getMessages(agent, prompt, input);
109
43
 
110
- const result = await agent.generateText({
111
- toolChoice: 'required',
112
- ...input,
113
- prompt,
44
+ const model = input.model ? agent.wrap(input.model) : agent.model;
45
+
46
+ const {
47
+ state,
48
+ machine,
49
+ previousPlan,
50
+ events,
51
+ goal,
52
+ model: _,
53
+ ...rest
54
+ } = input;
55
+
56
+ const machineState = input.machine
57
+ ? input.machine.resolveState({
58
+ ...input.state,
59
+ context: input.state.context,
60
+ })
61
+ : undefined;
62
+
63
+ const result = await generateText({
64
+ ...rest,
65
+ model,
114
66
  messages,
115
- tools: toolMap,
67
+ tools: toolMap as any,
68
+ toolChoice: input.toolChoice ?? 'required',
69
+ });
70
+
71
+ result.responseMessages.forEach((m) => {
72
+ const message: CoreMessage = m;
73
+
74
+ agent.addMessage({
75
+ ...message,
76
+ id: randomId(),
77
+ timestamp: Date.now(),
78
+ });
116
79
  });
117
80
 
118
81
  const singleResult = result.toolResults[0];
@@ -124,16 +87,25 @@ export async function simplePlanner<T extends AnyAgent>(
124
87
  }
125
88
 
126
89
  return {
90
+ planner: 'simple',
127
91
  goal: input.goal,
128
- state: input.state,
129
- execute: async (state) => {
130
- if (JSON.stringify(state) === JSON.stringify(input.state)) {
131
- return singleResult.result;
132
- }
133
- return undefined;
134
- },
92
+ goalState: input.state,
135
93
  nextEvent: singleResult.result,
136
- sessionId: agent.sessionId,
94
+ episodeId: agent.episodeId,
137
95
  timestamp: Date.now(),
96
+ paths: [
97
+ {
98
+ state: undefined,
99
+ steps: [
100
+ {
101
+ event: singleResult.result,
102
+ state:
103
+ machine && machineState
104
+ ? getNextSnapshot(machine, machineState, singleResult.result)
105
+ : undefined,
106
+ },
107
+ ],
108
+ },
109
+ ],
138
110
  };
139
111
  }
@@ -67,8 +67,8 @@ export const chainOfNote = setup({
67
67
  },
68
68
  }).createMachine({
69
69
  initial: 'searching',
70
- context: (x) => ({
71
- ...x.input,
70
+ context: ({ input }) => ({
71
+ ...input,
72
72
  searchResults: null,
73
73
  summaries: null,
74
74
  }),
@@ -76,8 +76,8 @@ export const chainOfNote = setup({
76
76
  searching: {
77
77
  invoke: {
78
78
  src: 'searchWiki',
79
- input: (x) => ({
80
- query: x.context.prompt,
79
+ input: ({ context }) => ({
80
+ query: context.prompt,
81
81
  }),
82
82
  onDone: {
83
83
  actions: assign({
@@ -90,8 +90,8 @@ export const chainOfNote = setup({
90
90
  extracting: {
91
91
  invoke: {
92
92
  src: 'extractSummaries',
93
- input: (x) => ({
94
- searchResult: x.context.searchResults!,
93
+ input: ({ context }) => ({
94
+ searchResult: context.searchResults!,
95
95
  }),
96
96
  onDone: {
97
97
  actions: assign({
@@ -104,52 +104,3 @@ export const chainOfNote = setup({
104
104
  generating: {},
105
105
  },
106
106
  });
107
-
108
- // export function chainOfNote() {
109
- // return {
110
- // generateText: async (x) => {
111
- // const passages = await wiki.search(x.prompt!, {
112
- // limit: 5,
113
- // });
114
-
115
- // const extracts = await Promise.all(
116
- // passages.results.map(async (p) => {
117
- // const summary = await wiki.summary(p.title);
118
- // return summary.extract;
119
- // })
120
- // );
121
- // x.agent?.addMessage({
122
- // content: x.prompt!,
123
- // id: Date.now() + '',
124
- // role: 'user',
125
- // timestamp: Date.now(),
126
- // });
127
- // const result = await generateText({
128
- // model: x.model,
129
- // system: `Task Description:
130
-
131
- // 1. Read the given question and five Wikipedia passages to gather relevant information.
132
-
133
- // 2. Write reading notes summarizing the key points from these passages.
134
-
135
- // 3. Discuss the relevance of the given question and Wikipedia passages.
136
-
137
- // 4. If some passages are relevant to the given question, provide a brief answer based on the passages.
138
-
139
- // 5. If no passage is relevant, direcly provide answer without considering the passages.
140
-
141
- // Passages: \n${extracts.join('\n')}`,
142
- // prompt: `${x.prompt!}`,
143
- // });
144
-
145
- // x.agent?.addMessage({
146
- // content: result.text,
147
- // id: Date.now() + '',
148
- // role: 'user',
149
- // timestamp: Date.now(),
150
- // });
151
-
152
- // return result;
153
- // },
154
- // } satisfies AgentStrategy;
155
- // }
package/src/text.ts CHANGED
@@ -1,9 +1,13 @@
1
- import type { CoreMessage, CoreTool, GenerateTextResult } from 'ai';
1
+ import {
2
+ generateText,
3
+ streamText,
4
+ type CoreMessage,
5
+ type CoreTool,
6
+ type GenerateTextResult,
7
+ } from 'ai';
2
8
  import {
3
9
  AgentGenerateTextOptions,
4
- AgentGenerateTextResult,
5
10
  AgentStreamTextOptions,
6
- AgentStreamTextResult,
7
11
  AnyAgent,
8
12
  } from './types';
9
13
  import { defaultTextTemplate } from './templates/defaultText';
@@ -15,7 +19,6 @@ import {
15
19
  fromPromise,
16
20
  toObserver,
17
21
  } from 'xstate';
18
- import { randomId } from './utils';
19
22
 
20
23
  /**
21
24
  * Gets an array of messages from the given prompt, based on the agent and options.
@@ -45,158 +48,39 @@ export async function getMessages(
45
48
  return messages;
46
49
  }
47
50
 
48
- export async function agentGenerateText<T extends AnyAgent>(
49
- agent: T,
50
- options: AgentGenerateTextOptions
51
- ): Promise<AgentGenerateTextResult> {
52
- const resolvedOptions = {
53
- ...agent.defaultOptions,
54
- ...options,
55
- correlationId: options.correlationId ?? randomId(),
56
- };
57
- // Generate a correlation ID if one is not provided
58
- const template = resolvedOptions.template ?? defaultTextTemplate;
59
- // TODO: check if messages was provided instead
60
- const id = randomId();
61
- const goal =
62
- typeof resolvedOptions.prompt === 'string'
63
- ? resolvedOptions.prompt
64
- : await resolvedOptions.prompt(agent);
65
-
66
- const promptWithContext = template({
67
- goal,
68
- context: resolvedOptions.context,
69
- });
70
-
71
- const messages = await getMessages(agent, promptWithContext, resolvedOptions);
72
-
73
- agent.addMessage({
74
- id,
75
- role: 'user',
76
- content: promptWithContext,
77
- timestamp: Date.now(),
78
- correlationId: resolvedOptions.correlationId,
79
- parentCorrelationId: resolvedOptions.parentCorrelationId,
80
- });
81
-
82
- const result = await agent.adapter.generateText({
83
- ...resolvedOptions,
84
- prompt: undefined,
85
- messages,
86
- });
87
-
88
- agent.addMessage({
89
- content: result.text,
90
- id,
91
- role: 'assistant',
92
- timestamp: Date.now(),
93
- responseId: id,
94
- result,
95
- correlationId: resolvedOptions.correlationId,
96
- parentCorrelationId: resolvedOptions.parentCorrelationId,
97
- });
98
-
99
- return {
100
- ...result,
101
- parentCorrelationId: resolvedOptions.parentCorrelationId,
102
- correlationId: resolvedOptions.correlationId,
103
- };
104
- }
105
-
106
- export async function agentStreamText(
107
- agent: AnyAgent,
108
- options: AgentStreamTextOptions
109
- ): Promise<AgentStreamTextResult> {
110
- const resolvedOptions = {
111
- ...agent.defaultOptions,
112
- ...options,
113
- correlationId: options.correlationId ?? randomId(),
114
- };
115
- const template = resolvedOptions.template ?? defaultTextTemplate;
116
-
117
- const id = randomId();
118
- const goal =
119
- typeof resolvedOptions.prompt === 'string'
120
- ? resolvedOptions.prompt
121
- : await resolvedOptions.prompt(agent);
122
-
123
- const promptWithContext = template({
124
- goal,
125
- context: resolvedOptions.context,
126
- });
127
-
128
- const messages = await getMessages(agent, promptWithContext, resolvedOptions);
129
-
130
- agent.addMessage({
131
- role: 'user',
132
- content: promptWithContext,
133
- id,
134
- timestamp: Date.now(),
135
- correlationId: resolvedOptions.correlationId,
136
- parentCorrelationId: resolvedOptions.parentCorrelationId,
137
- });
138
-
139
- const result = await agent.adapter.streamText({
140
- ...resolvedOptions,
141
- prompt: undefined,
142
- messages,
143
- onFinish: async (res) => {
144
- agent.addMessage({
145
- role: 'assistant',
146
- result: {
147
- text: res.text,
148
- finishReason: res.finishReason,
149
- logprobs: undefined,
150
- responseMessages: [],
151
- toolCalls: [],
152
- toolResults: [],
153
- usage: res.usage,
154
- warnings: res.warnings,
155
- rawResponse: res.rawResponse,
156
- roundtrips: [], // TODO: how do we get this information?,
157
- steps: res.steps,
158
- response: res.response,
159
- experimental_providerMetadata: res.experimental_providerMetadata,
160
- },
161
- content: res.text,
162
- id: randomId(),
163
- timestamp: Date.now(),
164
- responseId: id,
165
- correlationId: resolvedOptions.correlationId,
166
- parentCorrelationId: resolvedOptions.parentCorrelationId,
167
- });
168
- },
169
- });
170
-
171
- return {
172
- ...result,
173
- textStream: result.textStream,
174
- fullStream: result.fullStream,
175
- parentCorrelationId: resolvedOptions.parentCorrelationId,
176
- correlationId: resolvedOptions.correlationId,
177
- } as unknown as AgentStreamTextResult; // TODO: fix
178
- }
179
-
180
51
  export function fromTextStream<T extends AnyAgent>(
181
52
  agent: T,
182
- defaultOptions?: AgentStreamTextOptions
53
+ options?: AgentStreamTextOptions
183
54
  ): ObservableActorLogic<
184
55
  { textDelta: string },
185
56
  Omit<AgentStreamTextOptions, 'context'> & {
186
57
  context?: AgentStreamTextOptions['context'];
187
58
  }
188
59
  > {
60
+ const template = options?.template ?? defaultTextTemplate;
189
61
  return fromObservable(({ input }) => {
190
62
  const observers = new Set<Observer<{ textDelta: string }>>();
191
63
 
192
64
  // TODO: check if messages was provided instead
193
65
 
194
66
  (async () => {
195
- const result = await agentStreamText(agent, {
196
- ...defaultOptions,
197
- ...input,
67
+ const model = input.model ? agent.wrap(input.model) : agent.model;
68
+ const goal =
69
+ typeof input.prompt === 'string'
70
+ ? input.prompt
71
+ : await input.prompt(agent);
72
+ const promptWithContext = template({
73
+ goal,
198
74
  context: input.context,
199
75
  });
76
+ const messages = await getMessages(agent, promptWithContext, input);
77
+ const result = await streamText({
78
+ ...options,
79
+ ...input,
80
+ prompt: undefined, // overwritten by messages
81
+ model,
82
+ messages,
83
+ });
200
84
 
201
85
  for await (const part of result.fullStream) {
202
86
  if (part.type === 'text-delta') {
@@ -224,18 +108,41 @@ export function fromTextStream<T extends AnyAgent>(
224
108
 
225
109
  export function fromText<T extends AnyAgent>(
226
110
  agent: T,
227
- defaultOptions?: AgentGenerateTextOptions
111
+ options?: AgentGenerateTextOptions
228
112
  ): PromiseActorLogic<
229
113
  GenerateTextResult<Record<string, CoreTool<any, any>>>,
230
114
  Omit<AgentGenerateTextOptions, 'context'> & {
231
115
  context?: AgentGenerateTextOptions['context'];
232
116
  }
233
117
  > {
118
+ const resolvedOptions = {
119
+ ...agent.defaultOptions,
120
+ ...options,
121
+ };
122
+
123
+ const template = resolvedOptions.template ?? defaultTextTemplate;
124
+
234
125
  return fromPromise(async ({ input }) => {
235
- return await agentGenerateText(agent, {
236
- ...input,
237
- ...defaultOptions,
126
+ const goal =
127
+ typeof input.prompt === 'string'
128
+ ? input.prompt
129
+ : await input.prompt(agent);
130
+
131
+ const promptWithContext = template({
132
+ goal,
238
133
  context: input.context,
239
134
  });
135
+
136
+ const messages = await getMessages(agent, promptWithContext, input);
137
+
138
+ const model = input.model ? agent.wrap(input.model) : agent.model;
139
+
140
+ return await generateText({
141
+ ...input,
142
+ ...options,
143
+ prompt: undefined,
144
+ messages,
145
+ model,
146
+ });
240
147
  });
241
148
  }