Skip to main content

runifold_agent/agent/
execution.rs

1//! Canonical Agent execution engine and its private runtime helpers.
2
3use 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    /// Runs a user text turn inside an existing runtime context.
18    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    /// Streams real-time events while driving the canonical Agent loop.
32    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    /// Runs with write-ahead checkpoint persistence.
45    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    /// Resumes a persisted Agent execution.
62    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}