1use serde::{Deserialize, Serialize};
8use serde_json::{Value, json};
9use std::collections::BTreeMap;
10
11use crate::error::{AgentLoopError, Result};
12
13#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
15#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
16#[cfg_attr(feature = "openapi", schema(example = json!({"type":"function_call","call_id":"call_lookup_1","name":"lookup","arguments":"{\"query\":\"weather in Paris\"}","async":true})))]
17#[serde(tag = "type")]
18pub enum NativeToolCall {
19 #[serde(rename = "function_call")]
20 Function {
21 #[cfg_attr(feature = "openapi", schema(example = "call_lookup_1"))]
23 call_id: String,
24 #[cfg_attr(feature = "openapi", schema(example = "lookup"))]
26 name: String,
27 #[cfg_attr(
29 feature = "openapi",
30 schema(example = r#"{"query":"weather in Paris"}"#)
31 )]
32 arguments: String,
33 #[cfg_attr(feature = "openapi", schema(example = true))]
35 #[serde(rename = "async", default)]
36 asynchronous: bool,
37 },
38 #[serde(rename = "custom_tool_call")]
39 Custom {
40 #[cfg_attr(feature = "openapi", schema(example = "call_lookup_1"))]
42 call_id: String,
43 #[cfg_attr(feature = "openapi", schema(example = "lookup"))]
45 name: String,
46 #[cfg_attr(feature = "openapi", schema(example = "weather in Paris"))]
48 input: String,
49 #[cfg_attr(feature = "openapi", schema(example = true))]
51 #[serde(rename = "async", default)]
52 asynchronous: bool,
53 },
54}
55
56impl NativeToolCall {
57 pub fn id(&self) -> &str {
58 match self {
59 Self::Function { call_id, .. } | Self::Custom { call_id, .. } => call_id,
60 }
61 }
62 pub fn name(&self) -> &str {
63 match self {
64 Self::Function { name, .. } | Self::Custom { name, .. } => name,
65 }
66 }
67 pub fn is_async(&self) -> bool {
68 match self {
69 Self::Function { asynchronous, .. } | Self::Custom { asynchronous, .. } => {
70 *asynchronous
71 }
72 }
73 }
74 pub fn output(&self, output: &str) -> Value {
75 json!({"type": match self { Self::Function { .. } => "function_call_output", Self::Custom { .. } => "custom_tool_call_output" }, "call_id": self.id(), "output": output})
76 }
77 pub fn validate(&self) -> Result<()> {
78 if self.id().is_empty() || self.name().is_empty() {
79 return Err(AgentLoopError::llm(
80 "native tool call is missing its call_id or name",
81 ));
82 }
83 if let Self::Function { arguments, .. } = self {
84 let _: Value = serde_json::from_str(arguments).map_err(|_| {
85 AgentLoopError::llm("completed function call has invalid JSON arguments")
86 })?;
87 }
88 Ok(())
89 }
90}
91
92#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
93pub enum PendingCallState {
94 Queued,
95 Running,
96 Ready { output: String },
97 Delivered,
98}
99
100#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
101pub struct PendingCall {
102 pub call: NativeToolCall,
103 pub replay_safe: bool,
105 pub state: PendingCallState,
106}
107
108#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
109pub struct Delivery {
110 pub previous_response_id: String,
111 pub call_ids: Vec<String>,
112 pub input: Vec<Value>,
113}
114
115#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
117pub struct NativeAsyncCheckpoint {
118 pub latest_response_id: Option<String>,
119 #[serde(default)]
121 pub transcript_message_id: Option<String>,
122 #[serde(default)]
124 pub host_outcome: Option<Value>,
125 #[serde(default)]
126 pub completed_responses: u32,
127 #[serde(default)]
129 pub host_responses: Vec<Value>,
130 #[serde(default)]
131 pub response_in_flight: bool,
132 pub calls: BTreeMap<String, PendingCall>,
133 pub order: Vec<String>,
135 pub delivery: Option<Delivery>,
138 #[serde(default, skip_serializing_if = "Option::is_none")]
141 pub background_response: Option<crate::background_call::BackgroundResponseRecord>,
142}
143
144impl NativeAsyncCheckpoint {
145 pub fn register(&mut self, call: NativeToolCall, replay_safe: bool) -> Result<bool> {
148 call.validate()?;
149 if let Some(existing) = self.calls.get(call.id()) {
150 if existing.call != call || existing.replay_safe != replay_safe {
151 return Err(AgentLoopError::config(
152 "native call_id reused with different inputs or policy",
153 ));
154 }
155 return Ok(false);
156 }
157 self.order.push(call.id().to_owned());
158 self.calls.insert(
159 call.id().to_owned(),
160 PendingCall {
161 call,
162 replay_safe,
163 state: PendingCallState::Queued,
164 },
165 );
166 Ok(true)
167 }
168
169 pub fn start(&mut self, id: &str) -> Result<()> {
170 let pending = self
171 .calls
172 .get_mut(id)
173 .ok_or_else(|| AgentLoopError::config("unknown native call_id"))?;
174 if pending.state != PendingCallState::Queued {
175 return Err(AgentLoopError::config("native call is not queued"));
176 }
177 pending.state = PendingCallState::Running;
178 Ok(())
179 }
180
181 pub fn settle(&mut self, id: &str, output: String) -> Result<bool> {
183 let pending = self
184 .calls
185 .get_mut(id)
186 .ok_or_else(|| AgentLoopError::config("unknown native call_id"))?;
187 if matches!(
188 pending.state,
189 PendingCallState::Ready { .. } | PendingCallState::Delivered
190 ) {
191 return Ok(false);
192 }
193 pending.state = PendingCallState::Ready { output };
194 Ok(true)
195 }
196
197 pub fn response_completed(&mut self, response_id: String) -> Result<()> {
198 if response_id.is_empty() || self.delivery.is_some() {
199 return Err(AgentLoopError::config(
200 "response ID missing or delivery requires acknowledgement",
201 ));
202 }
203 if self.latest_response_id.as_ref() != Some(&response_id) {
204 self.completed_responses = self.completed_responses.saturating_add(1);
205 }
206 self.latest_response_id = Some(response_id);
207 Ok(())
208 }
209
210 pub fn recover(&mut self) -> Result<()> {
212 if self.transcript_message_id.is_some() {
213 return Err(AgentLoopError::store(
214 "native response transcript requires reconciliation",
215 ));
216 }
217 let ordered: std::collections::BTreeSet<_> = self.order.iter().collect();
218 if ordered.len() != self.calls.len()
219 || self.order.len() != self.calls.len()
220 || self
221 .calls
222 .iter()
223 .any(|(id, pending)| id != pending.call.id() || !ordered.contains(id))
224 {
225 return Err(AgentLoopError::store(
226 "native checkpoint call registry is inconsistent",
227 ));
228 }
229 for pending in self.calls.values() {
230 pending.call.validate()?;
231 }
232
233 if self.delivery.is_some() || self.response_in_flight {
234 return Err(AgentLoopError::store(
235 "native result delivery is uncertain; reconcile the response receipt before resuming",
236 ));
237 }
238 for pending in self.calls.values_mut() {
239 if pending.state == PendingCallState::Running {
240 pending.state = if pending.replay_safe {
241 PendingCallState::Queued
242 } else {
243 PendingCallState::Ready { output: json!({"error":"interrupted; execution outcome is uncertain; do not retry automatically"}).to_string() }
244 };
245 }
246 }
247 Ok(())
248 }
249
250 pub fn cancel(&mut self) {
251 for pending in self.calls.values_mut() {
252 if matches!(
253 pending.state,
254 PendingCallState::Queued | PendingCallState::Running
255 ) {
256 pending.state = PendingCallState::Ready {
257 output: json!({"error":"cancelled"}).to_string(),
258 };
259 }
260 }
261 }
262
263 pub fn prepare_delivery(&mut self) -> Result<Option<&Delivery>> {
267 if self.response_in_flight {
268 return Err(AgentLoopError::store(
269 "cannot deliver outputs while the latest response is incomplete",
270 ));
271 }
272 if self.delivery.is_some() {
273 return Err(AgentLoopError::store("native delivery already in flight"));
274 }
275 let ready: Vec<_> = self
276 .calls
277 .iter()
278 .filter_map(|(id, pending)| {
279 if let PendingCallState::Ready { output } = &pending.state {
280 Some((id.clone(), pending.call.output(output)))
281 } else {
282 None
283 }
284 })
285 .collect();
286 if ready.is_empty() {
287 return Ok(None);
288 }
289 let previous_response_id = self.latest_response_id.clone().ok_or_else(|| {
290 AgentLoopError::store("cannot deliver native outputs before a response completes")
291 })?;
292 let (call_ids, input) = ready.into_iter().unzip();
293 self.delivery = Some(Delivery {
294 previous_response_id,
295 call_ids,
296 input,
297 });
298 Ok(self.delivery.as_ref())
299 }
300
301 #[expect(
304 clippy::expect_used,
305 reason = "Pending delivery call IDs come only from the registered call map"
306 )]
307 pub fn acknowledge_delivery(&mut self, response_id: String) -> Result<()> {
308 if response_id.is_empty() {
309 return Err(AgentLoopError::config("empty native delivery response ID"));
310 }
311 if self
312 .delivery
313 .as_ref()
314 .is_some_and(|delivery| delivery.previous_response_id == response_id)
315 {
316 return Err(AgentLoopError::store(
317 "native delivery receipt must identify a new response",
318 ));
319 }
320 let Some(delivery) = self.delivery.take() else {
321 if self.latest_response_id.as_ref() == Some(&response_id) {
322 return Ok(());
323 }
324 return Err(AgentLoopError::store(
325 "no native delivery awaiting acknowledgement",
326 ));
327 };
328 for id in delivery.call_ids {
329 self.calls
330 .get_mut(&id)
331 .expect("delivery references registered calls")
332 .state = PendingCallState::Delivered;
333 }
334 if self.latest_response_id.as_ref() != Some(&response_id) {
335 self.completed_responses = self.completed_responses.saturating_add(1);
336 }
337 self.latest_response_id = Some(response_id);
338 Ok(())
339 }
340
341 pub fn can_complete(&self) -> bool {
342 self.transcript_message_id.is_none()
343 && !self.response_in_flight
344 && self.delivery.is_none()
345 && self
346 .calls
347 .values()
348 .all(|pending| pending.state == PendingCallState::Delivered)
349 }
350}
351
352#[cfg(test)]
353mod tests {
354 use super::*;
355 fn call(id: &str) -> NativeToolCall {
356 NativeToolCall::Function {
357 call_id: id.into(),
358 name: "lookup".into(),
359 arguments: "{}".into(),
360 asynchronous: true,
361 }
362 }
363 #[test]
364 fn outputs_follow_latest_response_and_deduplicate() {
365 let mut state = NativeAsyncCheckpoint::default();
366 state.register(call("slow"), true).unwrap();
367 state.register(call("fast"), true).unwrap();
368 assert!(!state.register(call("fast"), true).unwrap());
369 state.start("slow").unwrap();
370 state.start("fast").unwrap();
371 state.response_completed("response_launch".into()).unwrap();
372 state
373 .response_completed("response_independent".into())
374 .unwrap();
375 state.settle("fast", "first".into()).unwrap();
376 assert!(!state.settle("fast", "duplicate".into()).unwrap());
377 let delivery = state.prepare_delivery().unwrap().unwrap();
378 assert_eq!(delivery.previous_response_id, "response_independent");
379 assert_eq!(
380 delivery.input,
381 vec![json!({"type":"function_call_output","call_id":"fast","output":"first"})]
382 );
383 assert!(!state.can_complete());
384 state.acknowledge_delivery("response_fast".into()).unwrap();
385 state.acknowledge_delivery("response_fast".into()).unwrap();
386 state.settle("slow", "last".into()).unwrap();
387 assert_eq!(
388 state
389 .prepare_delivery()
390 .unwrap()
391 .unwrap()
392 .previous_response_id,
393 "response_fast"
394 );
395 state.acknowledge_delivery("response_final".into()).unwrap();
396 assert!(state.can_complete());
397 assert!(!state.register(call("fast"), true).unwrap());
398 }
399 #[test]
400 fn recovery_replays_only_safe_work_and_preserves_cancellation() {
401 let mut state = NativeAsyncCheckpoint::default();
402 state.register(call("safe"), true).unwrap();
403 state.register(call("unsafe"), false).unwrap();
404 state.start("safe").unwrap();
405 state.start("unsafe").unwrap();
406 let mut recovered: NativeAsyncCheckpoint =
407 serde_json::from_str(&serde_json::to_string(&state).unwrap()).unwrap();
408 recovered.recover().unwrap();
409 assert_eq!(recovered.calls["safe"].state, PendingCallState::Queued);
410 assert!(matches!(
411 recovered.calls["unsafe"].state,
412 PendingCallState::Ready { .. }
413 ));
414 recovered.cancel();
415 assert!(!recovered.settle("safe", "too late".into()).unwrap());
416 assert!(!recovered.can_complete());
417 recovered.response_completed("r1".into()).unwrap();
418 recovered.prepare_delivery().unwrap();
419 let before = recovered.clone();
420 assert!(recovered.recover().is_err());
421 assert_eq!(recovered, before);
422 recovered
423 .acknowledge_delivery("reconciled_receipt".into())
424 .unwrap();
425 assert!(recovered.can_complete());
426 }
427 #[test]
428 fn custom_calls_retain_input_async_and_original_id() {
429 let call: NativeToolCall = serde_json::from_value(json!({"type":"custom_tool_call","call_id":"custom","name":"query","input":"a raw query","async":true})).unwrap();
430 assert!(call.is_async());
431 assert_eq!(call.output("result")["type"], "custom_tool_call_output");
432 assert_eq!(serde_json::to_value(&call).unwrap()["async"], true);
433 assert_eq!(call.output("result")["call_id"], "custom");
434 }
435 #[test]
436 fn reused_ids_and_invalid_calls_fail_closed() {
437 let mut state = NativeAsyncCheckpoint::default();
438 state.register(call("same"), true).unwrap();
439 let changed = NativeToolCall::Function {
440 call_id: "same".into(),
441 name: "lookup".into(),
442 arguments: "{\"different\":true}".into(),
443 asynchronous: true,
444 };
445 assert!(state.register(changed, true).is_err());
446 assert!(state.register(call(""), true).is_err());
447 let malformed = NativeToolCall::Function {
448 call_id: "invalid".into(),
449 name: "lookup".into(),
450 arguments: "{".into(),
451 asynchronous: true,
452 };
453 assert!(state.register(malformed, true).is_err());
454 state.response_in_flight = true;
455 assert!(state.recover().is_err());
456 assert!(!state.can_complete());
457 }
458}