1use super::checkpointing::{
4 AgentProgress, save_checkpoint, validate_exact_usage, validate_usage_floor,
5};
6use super::observability::{consume_budget, emit_usage, record_domain, terminal_event};
7use super::{
8 Agent, AgentCheckpoint, AgentCheckpointPhase, AgentCheckpointState, AgentError,
9 AgentEventStream, AgentFuture, AgentObserver, AgentOutcome, AgentStreamEvent, Arc,
10 BufferedObserver, CheckpointCursor, ContentPart, Either, EventId, Instant, LifecycleEvent,
11 Message, ModelCallContext, ModelError, ModelErrorKind, ModelRequest, ModelResponse,
12 ModelStreamAccumulator, NoopObserver, ResumePolicy, Role, RunContext, RunEventKind, StreamExt,
13 ToolCall, Usage, emit_agent_event, select,
14};
15
16impl Agent {
17 pub fn run<'a>(
19 &'a self,
20 input: impl Into<String> + Send + 'a,
21 run: &'a RunContext,
22 ) -> AgentFuture<'a, Result<AgentOutcome, AgentError>> {
23 let input = input.into();
24 let state = self.initial_state(input, run.root_run_id().to_string());
25 Box::pin(async move {
26 self.execute_state(state, run, None, Arc::new(NoopObserver))
27 .await
28 })
29 }
30
31 pub fn stream<'a>(
33 &'a self,
34 input: impl Into<String> + Send + 'a,
35 run: &'a RunContext,
36 ) -> AgentEventStream<'a> {
37 let state = self.initial_state(input.into(), run.root_run_id().to_string());
38 let observer = BufferedObserver::default();
39 let events = observer.events();
40 let execution = Box::pin(self.execute_state(state, run, None, Arc::new(observer)));
41 AgentEventStream::new(execution, events)
42 }
43
44 pub fn run_checkpointed<'a>(
46 &'a self,
47 input: impl Into<String> + Send + 'a,
48 run: &'a RunContext,
49 checkpoint: &'a AgentCheckpoint,
50 ) -> AgentFuture<'a, Result<AgentOutcome, AgentError>> {
51 let input = input.into();
52 Box::pin(async move {
53 let mut state = self.initial_state(input, checkpoint.id().to_string());
54 state.usage = run.budget().usage();
55 let mut cursor = CheckpointCursor::create(checkpoint, run, &state)?;
56 self.execute_state(state, run, Some(&mut cursor), Arc::new(NoopObserver))
57 .await
58 })
59 }
60
61 pub fn resume<'a>(
63 &'a self,
64 checkpoint: &'a AgentCheckpoint,
65 run: &'a RunContext,
66 policy: ResumePolicy,
67 ) -> AgentFuture<'a, Result<AgentOutcome, AgentError>> {
68 Box::pin(async move {
69 let (envelope, mut state) = checkpoint.load()?;
70 self.validate_checkpoint_identity(&state)?;
71 if let Some(outcome) = state.outcome() {
72 validate_exact_usage(state.usage, run.budget().usage())?;
73 return Ok(outcome);
74 }
75 if let AgentCheckpointPhase::TurnInFlight { turn } = state.phase {
76 if policy == ResumePolicy::RejectAmbiguous {
77 return Err(AgentError::AmbiguousCheckpoint { turn });
78 }
79 validate_usage_floor(state.usage, run.budget().usage())?;
80 state.usage = run.budget().usage();
81 state.phase = AgentCheckpointPhase::ReadyForTurn;
82 } else {
83 validate_exact_usage(state.usage, run.budget().usage())?;
84 }
85 let mut cursor = CheckpointCursor::loaded(checkpoint, envelope);
86 self.execute_state(state, run, Some(&mut cursor), Arc::new(NoopObserver))
87 .await
88 })
89 }
90
91 fn initial_state(&self, input: String, execution_id: String) -> AgentCheckpointState {
92 let mut transcript = self.instructions.clone();
93 transcript.push(Message::user(input));
94 AgentCheckpointState {
95 execution_id,
96 agent: self.name.clone(),
97 model: self.model_ref.clone(),
98 transcript,
99 turns: 0,
100 tool_calls: 0,
101 delegations: 0,
102 usage: Usage::default(),
103 phase: AgentCheckpointPhase::ReadyForTurn,
104 }
105 }
106
107 async fn execute_state(
108 &self,
109 state: AgentCheckpointState,
110 run: &RunContext,
111 checkpoint: Option<&mut CheckpointCursor>,
112 observer: Arc<dyn AgentObserver>,
113 ) -> Result<AgentOutcome, AgentError> {
114 let started = run
115 .record(
116 RunEventKind::Lifecycle(LifecycleEvent::Started),
117 run.caused_by(),
118 )?
119 .map(|event| event.meta.event_id);
120 emit_agent_event(
121 observer.as_ref(),
122 AgentStreamEvent::Started {
123 agent: self.name.clone(),
124 },
125 )
126 .await;
127 let result = self
128 .run_loop(state, run, started, checkpoint, observer.as_ref())
129 .await;
130 let terminal = terminal_event(&self.name, &result);
131 run.record(terminal, started)?;
132 if let Ok(outcome) = &result {
133 emit_agent_event(
134 observer.as_ref(),
135 AgentStreamEvent::Completed {
136 outcome: outcome.clone(),
137 },
138 )
139 .await;
140 }
141 result
142 }
143
144 async fn run_loop(
145 &self,
146 state: AgentCheckpointState,
147 run: &RunContext,
148 caused_by: Option<EventId>,
149 mut checkpoint: Option<&mut CheckpointCursor>,
150 observer: &dyn AgentObserver,
151 ) -> Result<AgentOutcome, AgentError> {
152 self.validate_config()?;
153 let mut progress = AgentProgress::from(state);
154
155 loop {
156 Self::check_lifecycle(run)?;
157 if progress.turns >= self.config.max_turns {
158 return Err(AgentError::MaxTurns {
159 max_turns: self.config.max_turns,
160 });
161 }
162 save_checkpoint(
163 &mut checkpoint,
164 &self.checkpoint_state(
165 &progress,
166 run,
167 AgentCheckpointPhase::TurnInFlight {
168 turn: progress.turns + 1,
169 },
170 ),
171 )?;
172 consume_budget(
173 run,
174 Usage {
175 turns: 1,
176 ..Usage::default()
177 },
178 caused_by,
179 )?;
180 progress.turns += 1;
181 emit_agent_event(
182 observer,
183 AgentStreamEvent::TurnStarted {
184 turn: progress.turns,
185 },
186 )
187 .await;
188 emit_usage(observer, run).await;
189 record_domain(
190 run,
191 "turn.started",
192 serde_json::json!({"agent": self.name, "turn": progress.turns}),
193 caused_by,
194 )?;
195
196 let response = self
197 .invoke_model(
198 &progress.transcript,
199 run,
200 progress.turns,
201 caused_by,
202 observer,
203 )
204 .await?;
205
206 let calls = tool_calls_from(&response.content);
207 let assistant = Message::new(Role::Assistant, response.content.clone())
208 .map_err(|error| AgentError::Protocol(error.to_string()))?;
209 progress.transcript.push(assistant);
210
211 if calls.is_empty() {
212 if matches!(
213 response.finish_reason,
214 runifold_model::FinishReason::ToolCalls
215 ) {
216 return Err(AgentError::Protocol(
217 "model stopped for tool calls without emitting a tool call".into(),
218 ));
219 }
220 save_checkpoint(
221 &mut checkpoint,
222 &self.checkpoint_state(
223 &progress,
224 run,
225 AgentCheckpointPhase::Completed {
226 response: Box::new(response.clone()),
227 },
228 ),
229 )?;
230 return Ok(progress.outcome(response, run.budget().usage()));
231 }
232
233 self.execute_calls(calls, run, caused_by, &mut progress, observer)
234 .await?;
235 save_checkpoint(
236 &mut checkpoint,
237 &self.checkpoint_state(&progress, run, AgentCheckpointPhase::ReadyForTurn),
238 )?;
239 }
240 }
241
242 async fn invoke_model(
243 &self,
244 transcript: &[Message],
245 run: &RunContext,
246 turn: u32,
247 caused_by: Option<EventId>,
248 observer: &dyn AgentObserver,
249 ) -> Result<ModelResponse, AgentError> {
250 record_domain(
251 run,
252 "model.started",
253 serde_json::json!({
254 "agent": self.name,
255 "turn": turn,
256 "provider": self.model_ref.provider,
257 "model": self.model_ref.name,
258 }),
259 caused_by,
260 )?;
261 let response = match self
262 .stream_model_response(self.request(transcript)?, run, turn, observer)
263 .await
264 {
265 Ok(response) => response,
266 Err(error) => {
267 record_domain(
268 run,
269 "model.failed",
270 serde_json::json!({
271 "agent": self.name,
272 "turn": turn,
273 "kind": format!("{:?}", error.kind),
274 }),
275 caused_by,
276 )?;
277 return Err(error.into());
278 }
279 };
280 record_domain(
281 run,
282 "model.completed",
283 serde_json::json!({
284 "agent": self.name,
285 "turn": turn,
286 "finish_reason": response.finish_reason,
287 "usage": response.usage,
288 }),
289 caused_by,
290 )?;
291 consume_budget(run, response.usage.into(), caused_by)?;
292 emit_usage(observer, run).await;
293 Ok(response)
294 }
295
296 async fn stream_model_response(
297 &self,
298 request: ModelRequest,
299 run: &RunContext,
300 turn: u32,
301 observer: &dyn AgentObserver,
302 ) -> Result<ModelResponse, ModelError> {
303 let context = ModelCallContext::for_run(run);
304 let cancellation = context.cancellation().clone();
305 let opening = self.model.stream(request, context);
306 let mut stream = match select(Box::pin(cancellation.cancelled()), Box::pin(opening)).await {
307 Either::Left(_) => return Err(cancelled_model_error()),
308 Either::Right((result, _)) => result?,
309 };
310 let mut accumulator = ModelStreamAccumulator::new();
311 loop {
312 let next = stream.next();
313 let event = match select(Box::pin(cancellation.cancelled()), Box::pin(next)).await {
314 Either::Left(_) => return Err(cancelled_model_error()),
315 Either::Right((Some(event), _)) => event?,
316 Either::Right((None, _)) => {
317 return Err(ModelError::local(
318 ModelErrorKind::Protocol,
319 "model stream ended before a terminal response event",
320 ));
321 }
322 };
323 let response = accumulator.push(event.clone())?;
324 emit_agent_event(observer, AgentStreamEvent::Model { turn, event }).await;
325 if let Some(response) = response {
326 return Ok(response);
327 }
328 }
329 }
330
331 fn validate_config(&self) -> Result<(), AgentError> {
332 if self.name.trim().is_empty() {
333 return Err(AgentError::InvalidConfig(
334 "agent name cannot be empty".into(),
335 ));
336 }
337 if self.config.max_turns == 0 {
338 return Err(AgentError::InvalidConfig(
339 "max_turns must be greater than zero".into(),
340 ));
341 }
342 if let Some(collision) = self
343 .agents
344 .model_specs()
345 .into_iter()
346 .find(|spec| self.tools.contains(&spec.name))
347 {
348 return Err(AgentError::InvalidConfig(format!(
349 "callable name `{}` is registered as both a tool and an agent",
350 collision.name
351 )));
352 }
353 Ok(())
354 }
355
356 fn validate_checkpoint_identity(&self, state: &AgentCheckpointState) -> Result<(), AgentError> {
357 if state.agent != self.name || state.model != self.model_ref {
358 return Err(runifold_core::CheckpointError::new(
359 runifold_core::CheckpointErrorKind::InvalidPayload,
360 "checkpoint Agent or model identity does not match",
361 )
362 .into());
363 }
364 Ok(())
365 }
366
367 fn checkpoint_state(
368 &self,
369 progress: &AgentProgress,
370 run: &RunContext,
371 phase: AgentCheckpointPhase,
372 ) -> AgentCheckpointState {
373 AgentCheckpointState {
374 execution_id: progress.execution_id.clone(),
375 agent: self.name.clone(),
376 model: self.model_ref.clone(),
377 transcript: progress.transcript.clone(),
378 turns: progress.turns,
379 tool_calls: progress.tool_calls,
380 delegations: progress.delegations,
381 usage: run.budget().usage(),
382 phase,
383 }
384 }
385
386 pub(super) fn check_lifecycle(run: &RunContext) -> Result<(), AgentError> {
387 let error = if run.cancellation().is_cancelled() {
388 Some((
389 runifold_model::ModelErrorKind::Cancelled,
390 "agent run was cancelled",
391 ))
392 } else if run
393 .deadline()
394 .is_some_and(|deadline| deadline <= Instant::now())
395 {
396 Some((
397 runifold_model::ModelErrorKind::DeadlineExceeded,
398 "agent run deadline elapsed",
399 ))
400 } else {
401 None
402 };
403 if let Some((kind, message)) = error {
404 return Err(runifold_model::ModelError::local(kind, message).into());
405 }
406 Ok(())
407 }
408
409 fn request(&self, transcript: &[Message]) -> Result<ModelRequest, AgentError> {
410 let (first, rest) = transcript
411 .split_first()
412 .ok_or_else(|| AgentError::Protocol("agent transcript is empty".into()))?;
413 let mut request = ModelRequest::new(self.model_ref.clone(), first.clone());
414 request.messages.extend_from_slice(rest);
415 request.tools = self.tools.model_specs();
416 request.tools.extend(self.agents.model_specs());
417 request.feature_policy = self.config.feature_policy;
418 request.output_format.clone_from(&self.output_format);
419 Ok(request)
420 }
421}
422
423fn cancelled_model_error() -> ModelError {
424 ModelError::local(ModelErrorKind::Cancelled, "model invocation was cancelled")
425}
426
427fn tool_calls_from(content: &[ContentPart]) -> Vec<ToolCall> {
428 content
429 .iter()
430 .filter_map(|part| match part {
431 ContentPart::ToolCall(call) => Some(call.clone()),
432 _ => None,
433 })
434 .collect()
435}