1mod approve;
11mod config;
12mod run;
13mod session;
14mod sink;
15mod terminal;
16
17pub use approve::{AllowAll, ApprovalRequest, Approver, Decision};
18pub use config::EngineConfig;
19pub use session::Session;
20pub use sink::{EventSink, FnSink, NullSink};
21pub use tokio_util::sync::CancellationToken;
24
25#[cfg(test)]
26mod tests {
27 #![allow(clippy::unnecessary_literal_bound)]
30
31 use super::*;
32 use async_trait::async_trait;
33 use locode_protocol::{
34 ContentBlock, Conversation, Event, Message, ReasoningFormat, Role, Status, Usage,
35 reconstruct_conversation,
36 };
37 use locode_provider::{
38 Completion, ConversationRequest, MockProvider, Provider, ProviderError, StopReason,
39 };
40 use locode_tools::{Registry, Tool, ToolCtx, ToolError, ToolKind, ToolOutput};
41 use serde::Serialize;
42 use serde_json::{Value, json};
43 use std::sync::{Arc, Mutex};
44 use std::time::Duration;
45
46 #[derive(Serialize)]
49 struct EchoOut {
50 echoed: String,
51 }
52 impl ToolOutput for EchoOut {
53 fn to_prompt_text(&self) -> String {
54 self.echoed.clone()
55 }
56 }
57
58 struct Echo;
59 #[async_trait]
60 impl Tool for Echo {
61 type Args = Value;
62 type Output = EchoOut;
63 fn kind(&self) -> ToolKind {
64 ToolKind::Shell
65 }
66 fn description(&self) -> &str {
67 "echo"
68 }
69 async fn run(&self, _ctx: &ToolCtx, args: Value) -> Result<EchoOut, ToolError> {
70 Ok(EchoOut {
71 echoed: args.to_string(),
72 })
73 }
74 }
75
76 struct Boom;
77 #[async_trait]
78 impl Tool for Boom {
79 type Args = Value;
80 type Output = EchoOut;
81 fn kind(&self) -> ToolKind {
82 ToolKind::Shell
83 }
84 fn description(&self) -> &str {
85 "boom"
86 }
87 async fn run(&self, _ctx: &ToolCtx, _args: Value) -> Result<EchoOut, ToolError> {
88 Err(ToolError::Fatal("boom aborted the turn".into()))
89 }
90 }
91
92 fn text_turn(text: &str) -> Completion {
95 Completion {
96 content: vec![ContentBlock::Text { text: text.into() }],
97 usage: Usage::default(),
98 stop: StopReason::EndTurn,
99 }
100 }
101
102 fn tool_turn(id: &str, name: &str) -> Completion {
103 Completion {
104 content: vec![ContentBlock::ToolUse {
105 id: id.into(),
106 name: name.into(),
107 input: json!({}),
108 }],
109 usage: Usage::default(),
110 stop: StopReason::ToolUse,
111 }
112 }
113
114 fn config() -> EngineConfig {
115 EngineConfig {
116 session_id: "sess-1".into(),
117 harness: "grok".into(),
118 api_schema: "mock".into(),
119 model: "mock-1".into(),
120 max_turns: None,
121 resample_retries: 2,
122 resample_backoff: Duration::ZERO, ..EngineConfig::default()
124 }
125 }
126
127 fn session_with(
129 script: Vec<Result<Completion, ProviderError>>,
130 registry: Registry,
131 cfg: EngineConfig,
132 ) -> (Session, Arc<Mutex<Vec<Event>>>) {
133 let events = Arc::new(Mutex::new(Vec::new()));
134 let sink_events = Arc::clone(&events);
135 let sink = Box::new(FnSink(move |event| {
136 sink_events.lock().unwrap().push(event);
137 }));
138 let provider = Arc::new(MockProvider::with_results(script));
139 let session = Session::new(provider, registry, vec![], cfg, sink);
140 (session, events)
141 }
142
143 fn echo_registry() -> Registry {
144 let mut reg = Registry::new();
145 reg.register("echo", Echo);
146 reg
147 }
148
149 fn dump(events: &Arc<Mutex<Vec<Event>>>) -> Vec<Event> {
150 events.lock().unwrap().clone()
151 }
152
153 #[tokio::test]
156 async fn completed_with_no_tools() {
157 let (mut s, events) =
158 session_with(vec![Ok(text_turn("all done"))], Registry::new(), config());
159 let report = s.run_text("hi").await;
160 assert_eq!(report.status, Status::Completed);
161 assert_eq!(report.final_message.as_deref(), Some("all done"));
162 assert_eq!(report.turns, 1);
163 assert!(report.tool_calls.is_empty());
164 assert_eq!(report.api_schema, "mock");
165 let evs = dump(&events);
167 assert!(matches!(evs.first(), Some(Event::Init { .. })));
168 assert!(matches!(evs.last(), Some(Event::Result { .. })));
169 }
170
171 #[tokio::test]
172 async fn tool_call_then_complete() {
173 let (mut s, _e) = session_with(
174 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
175 echo_registry(),
176 config(),
177 );
178 let report = s.run_text("go").await;
179 assert_eq!(report.status, Status::Completed);
180 assert_eq!(report.turns, 2);
181 assert_eq!(report.tool_calls.len(), 1);
182 assert!(report.tool_calls[0].ok);
183 assert_eq!(report.tool_calls[0].name, "echo");
184 }
185
186 #[tokio::test]
187 async fn hits_max_turns_after_dispatch() {
188 let mut cfg = config();
190 cfg.max_turns = Some(2);
191 let (mut s, _e) = session_with(
192 vec![
193 Ok(tool_turn("c1", "echo")),
194 Ok(tool_turn("c2", "echo")),
195 Ok(tool_turn("c3", "echo")),
196 ],
197 echo_registry(),
198 cfg,
199 );
200 let report = s.run_text("go").await;
201 assert_eq!(report.status, Status::MaxTurns);
202 assert_eq!(report.turns, 2);
203 assert_eq!(report.tool_calls.len(), 2);
204 }
205
206 #[tokio::test]
207 async fn model_error_after_bounded_retry() {
208 let script = vec![
210 Err(ProviderError::Transport("reset".into())),
211 Err(ProviderError::Transport("reset".into())),
212 Err(ProviderError::Transport("reset".into())),
213 ];
214 let (mut s, events) = session_with(script, Registry::new(), config());
215 let report = s.run_text("go").await;
216 assert_eq!(report.status, Status::ModelError);
217 assert!(report.error.is_some());
218 assert_eq!(report.turns, 0);
219 let retries = dump(&events)
221 .iter()
222 .filter(|e| matches!(e, Event::Error { .. }))
223 .count();
224 assert_eq!(retries, 2);
225 }
226
227 #[tokio::test]
228 async fn model_error_non_retryable_is_immediate() {
229 let (mut s, events) = session_with(
230 vec![Err(ProviderError::ContextOverflow)],
231 Registry::new(),
232 config(),
233 );
234 let report = s.run_text("go").await;
235 assert_eq!(report.status, Status::ModelError);
236 let retries = dump(&events)
237 .iter()
238 .filter(|e| matches!(e, Event::Error { .. }))
239 .count();
240 assert_eq!(retries, 0, "a non-retryable error must not resample");
241 }
242
243 #[tokio::test]
244 async fn fatal_tool_error_ends_the_run() {
245 let mut reg = Registry::new();
246 reg.register("boom", Boom);
247 let (mut s, _e) = session_with(vec![Ok(tool_turn("c1", "boom"))], reg, config());
248 let report = s.run_text("go").await;
249 assert_eq!(report.status, Status::Error);
250 assert!(report.error.is_some());
251 assert_eq!(report.tool_calls.len(), 1);
253 assert!(!report.tool_calls[0].ok);
254 }
255
256 #[tokio::test]
260 async fn empty_completion_resamples_then_succeeds() {
261 let empty = Completion {
262 content: vec![ContentBlock::Reasoning {
263 format: ReasoningFormat::Anthropic,
264 text: "thinking only".into(),
265 signature: Some("sig".into()),
266 payload: None,
267 }],
268 usage: Usage::default(),
269 stop: StopReason::MaxTokens,
270 };
271 let (mut session, _events) = session_with(
272 vec![Ok(empty), Ok(text_turn("recovered"))],
273 echo_registry(),
274 config(),
275 );
276 let report = session.run_text("go").await;
277 assert_eq!(report.status, Status::Completed);
278 assert_eq!(report.final_message.as_deref(), Some("recovered"));
279 assert_eq!(report.stop_reason.as_deref(), Some("end_turn"));
280 }
281
282 #[tokio::test]
283 async fn persistent_empty_completions_are_model_error() {
284 let empty = || Completion {
285 content: vec![],
286 usage: Usage::default(),
287 stop: StopReason::MaxTokens,
288 };
289 let (mut session, _events) = session_with(
291 vec![Ok(empty()), Ok(empty()), Ok(empty())],
292 echo_registry(),
293 config(),
294 );
295 let report = session.run_text("go").await;
296 assert_eq!(report.status, Status::ModelError);
297 assert!(
298 report
299 .error
300 .as_deref()
301 .unwrap_or("")
302 .contains("empty completion"),
303 "error names the cause: {:?}",
304 report.error
305 );
306 assert_eq!(report.stop_reason, None, "no completion was accepted");
307 }
308
309 #[tokio::test]
312 async fn mid_batch_abort_synthesizes_results() {
313 let mut reg = Registry::new();
316 reg.register("boom", Boom);
317 reg.register("echo", Echo);
318 let completion = Completion {
319 content: vec![
320 ContentBlock::ToolUse {
321 id: "c_boom".into(),
322 name: "boom".into(),
323 input: json!({}),
324 },
325 ContentBlock::ToolUse {
326 id: "c_echo".into(),
327 name: "echo".into(),
328 input: json!({}),
329 },
330 ],
331 usage: Usage::default(),
332 stop: StopReason::ToolUse,
333 };
334 let (mut s, events) = session_with(vec![Ok(completion)], reg, config());
335 let report = s.run_text("go").await;
336 assert_eq!(report.status, Status::Error);
337
338 let evs = dump(&events);
340 let answered: Vec<String> = evs
341 .iter()
342 .filter_map(|e| match e {
343 Event::Message { message } if message.role == Role::User => Some(&message.content),
344 _ => None,
345 })
346 .flatten()
347 .filter_map(|b| match b {
348 ContentBlock::ToolResult { tool_use_id, .. } => Some(tool_use_id.clone()),
349 _ => None,
350 })
351 .collect();
352 assert!(answered.iter().any(|id| id == "c_boom"));
353 assert!(
354 answered.iter().any(|id| id == "c_echo"),
355 "the un-run echo must be paired"
356 );
357 assert_eq!(report.tool_calls.len(), 1);
359 }
360
361 #[tokio::test]
364 async fn thinking_block_is_appended_verbatim() {
365 let completion = Completion {
366 content: vec![
367 ContentBlock::Reasoning {
368 format: ReasoningFormat::Anthropic,
369 text: "reasoning".into(),
370 signature: Some("sig-xyz".into()),
371 payload: None,
372 },
373 ContentBlock::Text {
374 text: "answer".into(),
375 },
376 ],
377 usage: Usage::default(),
378 stop: StopReason::EndTurn,
379 };
380 let (mut s, events) = session_with(vec![Ok(completion)], Registry::new(), config());
381 let report = s.run_text("think").await;
382 assert_eq!(report.status, Status::Completed);
383 assert_eq!(report.final_message.as_deref(), Some("answer"));
384 let has_thinking = dump(&events).iter().any(|e| match e {
386 Event::Message { message } if message.role == Role::Assistant => {
387 message.content.iter().any(|b| {
388 matches!(
389 b,
390 ContentBlock::Reasoning { signature: Some(sig), .. } if sig == "sig-xyz"
391 )
392 })
393 }
394 _ => false,
395 });
396 assert!(
397 has_thinking,
398 "thinking + signature must survive into history"
399 );
400 }
401
402 #[tokio::test]
403 async fn events_reconstruct_the_history() {
404 let (mut s, events) = session_with(
405 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
406 echo_registry(),
407 config(),
408 );
409 let _ = s.run_text("go").await;
410 let rebuilt: Conversation = reconstruct_conversation(&dump(&events));
411 let roles: Vec<Role> = rebuilt.messages.iter().map(|m| m.role).collect();
413 assert_eq!(
414 roles,
415 vec![Role::User, Role::Assistant, Role::User, Role::Assistant]
416 );
417 }
418
419 use std::sync::atomic::{AtomicUsize, Ordering};
422
423 struct Counting(Arc<AtomicUsize>);
425 #[async_trait]
426 impl Tool for Counting {
427 type Args = Value;
428 type Output = EchoOut;
429 fn kind(&self) -> ToolKind {
430 ToolKind::Shell
431 }
432 fn description(&self) -> &str {
433 "counting"
434 }
435 async fn run(&self, _ctx: &ToolCtx, _args: Value) -> Result<EchoOut, ToolError> {
436 self.0.fetch_add(1, Ordering::SeqCst);
437 Ok(EchoOut {
438 echoed: "ran".into(),
439 })
440 }
441 }
442
443 type SeenKinds = Arc<Mutex<Vec<(String, Option<ToolKind>)>>>;
444
445 struct DenyNamed {
448 deny: Vec<&'static str>,
449 seen_kinds: SeenKinds,
450 }
451 #[async_trait]
452 impl Approver for DenyNamed {
453 async fn decide(&self, request: &ApprovalRequest<'_>) -> Decision {
454 self.seen_kinds
455 .lock()
456 .unwrap()
457 .push((request.tool_name.to_owned(), request.kind));
458 if self.deny.contains(&request.tool_name) {
459 Decision::Deny {
460 reason: format!("{} is not allowed here", request.tool_name),
461 }
462 } else {
463 Decision::Allow
464 }
465 }
466 }
467
468 fn approvals(events: &Arc<Mutex<Vec<Event>>>) -> Vec<(String, String, String)> {
469 dump(events)
470 .iter()
471 .filter_map(|e| match e {
472 Event::Approval {
473 tool_use_id,
474 tool_name,
475 decision,
476 ..
477 } => Some((tool_use_id.clone(), tool_name.clone(), decision.clone())),
478 _ => None,
479 })
480 .collect()
481 }
482
483 #[tokio::test]
484 async fn deny_is_a_soft_paired_error_and_the_run_continues() {
485 let ran = Arc::new(AtomicUsize::new(0));
486 let mut reg = Registry::new();
487 reg.register("counting", Counting(Arc::clone(&ran)));
488 let (s, events) = session_with(
489 vec![Ok(tool_turn("c1", "counting")), Ok(text_turn("done"))],
490 reg,
491 config(),
492 );
493 let seen = Arc::new(Mutex::new(Vec::new()));
494 let mut s = s.with_approver(Arc::new(DenyNamed {
495 deny: vec!["counting"],
496 seen_kinds: Arc::clone(&seen),
497 }));
498 let report = s.run_text("go").await;
499
500 assert_eq!(report.status, Status::Completed);
502 assert_eq!(ran.load(Ordering::SeqCst), 0, "denied tool must not run");
503
504 assert_eq!(report.tool_calls.len(), 1);
506 let record = &report.tool_calls[0];
507 assert!(!record.ok);
508 assert_eq!(
509 record.denial_reason.as_deref(),
510 Some("counting is not allowed here")
511 );
512 assert_eq!(record.kind, "shell", "kind still recorded on denial");
513
514 let denied_result = dump(&events).iter().any(|e| match e {
516 Event::Message { message } => message.content.iter().any(|b| {
517 matches!(
518 b,
519 ContentBlock::ToolResult { tool_use_id, is_error: true, content, .. }
520 if tool_use_id == "c1"
521 && content.iter().any(|c| matches!(
522 c,
523 locode_protocol::ResultChunk::Text { text }
524 if text == "tool call denied: counting is not allowed here"
525 ))
526 )
527 }),
528 _ => false,
529 });
530 assert!(denied_result, "the model sees the denial reason, paired");
531
532 assert_eq!(
534 approvals(&events),
535 vec![("c1".into(), "counting".into(), "deny".into())]
536 );
537 }
538
539 #[tokio::test]
540 async fn deny_then_allow_within_one_batch_keeps_order_and_pairing() {
541 let ran = Arc::new(AtomicUsize::new(0));
542 let mut reg = Registry::new();
543 reg.register("blocked", Counting(Arc::clone(&ran)));
544 reg.register("echo", Echo);
545 let batch = Completion {
546 content: vec![
547 ContentBlock::ToolUse {
548 id: "c1".into(),
549 name: "blocked".into(),
550 input: json!({}),
551 },
552 ContentBlock::ToolUse {
553 id: "c2".into(),
554 name: "echo".into(),
555 input: json!({}),
556 },
557 ],
558 usage: Usage::default(),
559 stop: StopReason::ToolUse,
560 };
561 let (s, events) = session_with(vec![Ok(batch), Ok(text_turn("done"))], reg, config());
562 let mut s = s.with_approver(Arc::new(DenyNamed {
563 deny: vec!["blocked"],
564 seen_kinds: Arc::new(Mutex::new(Vec::new())),
565 }));
566 let report = s.run_text("go").await;
567 assert_eq!(report.status, Status::Completed);
568 assert_eq!(ran.load(Ordering::SeqCst), 0);
569
570 let pairs: Vec<(String, bool)> = dump(&events)
572 .iter()
573 .filter_map(|e| match e {
574 Event::Message { message } if message.role == Role::User => Some(&message.content),
575 _ => None,
576 })
577 .flatten()
578 .filter_map(|b| match b {
579 ContentBlock::ToolResult {
580 tool_use_id,
581 is_error,
582 ..
583 } => Some((tool_use_id.clone(), *is_error)),
584 _ => None,
585 })
586 .collect();
587 assert_eq!(pairs, vec![("c1".into(), true), ("c2".into(), false)]);
588
589 assert_eq!(report.tool_calls.len(), 2);
591 assert!(report.tool_calls[0].denial_reason.is_some());
592 assert_eq!(report.tool_calls[0].kind, "shell");
593 assert!(report.tool_calls[1].ok);
594 assert_eq!(report.tool_calls[1].denial_reason, None);
595
596 assert_eq!(
598 approvals(&events),
599 vec![
600 ("c1".into(), "blocked".into(), "deny".into()),
601 ("c2".into(), "echo".into(), "allow".into()),
602 ]
603 );
604 }
605
606 #[tokio::test]
607 async fn approval_request_carries_the_registry_kind() {
608 let seen = Arc::new(Mutex::new(Vec::new()));
609 let (s, _e) = session_with(
610 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
611 echo_registry(),
612 config(),
613 );
614 let mut s = s.with_approver(Arc::new(DenyNamed {
615 deny: vec![],
616 seen_kinds: Arc::clone(&seen),
617 }));
618 let _ = s.run_text("go").await;
619 let seen = seen.lock().unwrap();
620 assert_eq!(seen.len(), 1);
621 assert_eq!(seen[0].0, "echo");
622 assert_eq!(
623 seen[0].1,
624 Some(ToolKind::Shell),
625 "kind resolves from the registry pre-dispatch"
626 );
627 }
628
629 #[tokio::test]
633 async fn async_approver_suspends_the_call_until_resolved() {
634 struct OneshotApprover(Mutex<Option<tokio::sync::oneshot::Receiver<Decision>>>);
635 #[async_trait]
636 impl Approver for OneshotApprover {
637 async fn decide(&self, _request: &ApprovalRequest<'_>) -> Decision {
638 let rx = self.0.lock().unwrap().take().expect("one decision");
639 rx.await.expect("decider dropped")
640 }
641 }
642
643 let (tx, rx) = tokio::sync::oneshot::channel();
644 let (s, _e) = session_with(
645 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
646 echo_registry(),
647 config(),
648 );
649 let mut s = s.with_approver(Arc::new(OneshotApprover(Mutex::new(Some(rx)))));
650
651 let ui = tokio::spawn(async move {
653 tokio::task::yield_now().await;
654 let _ = tx.send(Decision::Allow);
655 });
656 let report = s.run_text("go").await;
657 ui.await.expect("ui task");
658 assert_eq!(report.status, Status::Completed);
659 assert_eq!(report.tool_calls.len(), 1);
660 assert!(report.tool_calls[0].ok);
661 }
662
663 #[tokio::test]
664 async fn allowed_calls_emit_approval_events_by_default() {
665 let (mut s, events) = session_with(
667 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
668 echo_registry(),
669 config(),
670 );
671 let report = s.run_text("go").await;
672 assert_eq!(report.status, Status::Completed);
673 assert_eq!(
674 approvals(&events),
675 vec![("c1".into(), "echo".into(), "allow".into())]
676 );
677 assert_eq!(report.tool_calls[0].denial_reason, None);
679 }
680
681 struct HangingProvider;
686 #[async_trait]
687 impl Provider for HangingProvider {
688 #[allow(clippy::unnecessary_literal_bound)]
689 fn api_schema(&self) -> &str {
690 "mock"
691 }
692 async fn complete(
693 &self,
694 _request: &ConversationRequest,
695 ) -> Result<Completion, ProviderError> {
696 tokio::time::sleep(Duration::from_hours(1)).await;
697 Err(ProviderError::Transport("unreachable".into()))
698 }
699 }
700
701 struct WaitsForCancel;
704 #[async_trait]
705 impl Tool for WaitsForCancel {
706 type Args = Value;
707 type Output = EchoOut;
708 fn kind(&self) -> ToolKind {
709 ToolKind::Shell
710 }
711 fn description(&self) -> &str {
712 "waits"
713 }
714 async fn run(&self, ctx: &ToolCtx, _args: Value) -> Result<EchoOut, ToolError> {
715 ctx.cancel.cancelled().await;
716 Ok(EchoOut {
717 echoed: "stopped cooperatively".into(),
718 })
719 }
720 }
721
722 #[tokio::test]
723 async fn cancel_mid_sample_yields_cancelled_report() {
724 let events = Arc::new(Mutex::new(Vec::new()));
725 let sink_events = Arc::clone(&events);
726 let sink = Box::new(FnSink(move |event| {
727 sink_events.lock().unwrap().push(event);
728 }));
729 let mut s = Session::new(
730 Arc::new(HangingProvider),
731 Registry::new(),
732 vec![],
733 config(),
734 sink,
735 );
736 let handle = s.cancel_handle();
737 let canceller = tokio::spawn(async move {
738 tokio::time::sleep(Duration::from_millis(20)).await;
739 handle.cancel();
740 handle.cancel(); });
742 let report = s.run_text("go").await;
743 canceller.await.expect("canceller");
744
745 assert_eq!(report.status, Status::Cancelled);
746 assert_eq!(report.error, None, "cancelled is a stop, not a fault");
747 assert_eq!(report.final_message, None, "no assistant text this run");
748 assert_eq!(report.turns, 0, "no completion was accepted");
749 let roles: Vec<Role> = s.history().iter().map(|m| m.role).collect();
751 assert_eq!(roles, vec![Role::User]);
752 let evs = dump(&events);
754 assert!(
755 matches!(evs.last(), Some(Event::Result { report }) if report.status == Status::Cancelled)
756 );
757 }
758
759 #[tokio::test]
760 async fn cancel_mid_batch_pairs_the_rest_synthetically() {
761 let mut reg = Registry::new();
765 reg.register("waits", WaitsForCancel);
766 reg.register("echo", Echo);
767 let batch = Completion {
768 content: vec![
769 ContentBlock::ToolUse {
770 id: "c_wait".into(),
771 name: "waits".into(),
772 input: json!({}),
773 },
774 ContentBlock::ToolUse {
775 id: "c_echo".into(),
776 name: "echo".into(),
777 input: json!({}),
778 },
779 ],
780 usage: Usage::default(),
781 stop: StopReason::ToolUse,
782 };
783 let (s, events) = session_with(vec![Ok(batch)], reg, config());
784 let mut s = s; let handle = s.cancel_handle();
786 let canceller = tokio::spawn(async move {
787 tokio::time::sleep(Duration::from_millis(20)).await;
788 handle.cancel();
789 });
790 let report = s.run_text("go").await;
791 canceller.await.expect("canceller");
792
793 assert_eq!(report.status, Status::Cancelled);
794 assert_eq!(report.tool_calls.len(), 1);
797 assert_eq!(report.tool_calls[0].id, "c_wait");
798 assert!(report.tool_calls[0].ok);
799 assert_eq!(report.tool_calls[0].denial_reason, None);
800
801 let pairs: Vec<(String, bool)> = dump(&events)
803 .iter()
804 .filter_map(|e| match e {
805 Event::Message { message } if message.role == Role::User => Some(&message.content),
806 _ => None,
807 })
808 .flatten()
809 .filter_map(|b| match b {
810 ContentBlock::ToolResult {
811 tool_use_id,
812 is_error,
813 ..
814 } => Some((tool_use_id.clone(), *is_error)),
815 _ => None,
816 })
817 .collect();
818 assert_eq!(
819 pairs,
820 vec![("c_wait".into(), false), ("c_echo".into(), true)]
821 );
822 assert_eq!(
824 approvals(&events),
825 vec![("c_wait".into(), "waits".into(), "allow".into())]
826 );
827 }
828
829 #[tokio::test]
832 async fn cancelled_session_continues_on_the_next_run_with_a_fresh_token() {
833 let mut reg = Registry::new();
834 reg.register("waits", WaitsForCancel);
835 let (s, _e) = session_with(
836 vec![Ok(tool_turn("c1", "waits")), Ok(text_turn("second run"))],
837 reg,
838 config(),
839 );
840 let mut s = s;
841 let handle1 = s.cancel_handle();
842 let canceller = tokio::spawn(async move {
843 tokio::time::sleep(Duration::from_millis(20)).await;
844 handle1.cancel();
845 });
846 let r1 = s.run_text("q1").await;
847 canceller.await.expect("canceller");
848 assert_eq!(r1.status, Status::Cancelled);
849
850 assert!(!s.cancel_handle().is_cancelled());
853 let r2 = s.run_text("q2").await;
854 assert_eq!(r2.status, Status::Completed);
855 assert_eq!(r2.final_message.as_deref(), Some("second run"));
856 assert!(s.history().len() >= 4, "history: {:?}", s.history().len());
858 }
859
860 struct CapturingProvider {
865 inner: MockProvider,
866 requests: Arc<Mutex<Vec<Vec<Message>>>>,
867 }
868 #[async_trait]
869 impl Provider for CapturingProvider {
870 #[allow(clippy::unnecessary_literal_bound)]
871 fn api_schema(&self) -> &str {
872 "mock"
873 }
874 async fn complete(
875 &self,
876 request: &ConversationRequest,
877 ) -> Result<Completion, ProviderError> {
878 self.requests.lock().unwrap().push(request.messages.clone());
879 self.inner.complete(request).await
880 }
881 }
882
883 #[allow(clippy::type_complexity)]
885 fn capturing_session_with(
886 script: Vec<Result<Completion, ProviderError>>,
887 registry: Registry,
888 ) -> (
889 Session,
890 Arc<Mutex<Vec<Vec<Message>>>>,
891 Arc<Mutex<Vec<Event>>>,
892 ) {
893 let requests = Arc::new(Mutex::new(Vec::new()));
894 let events = Arc::new(Mutex::new(Vec::new()));
895 let sink_events = Arc::clone(&events);
896 let sink = Box::new(FnSink(move |event| {
897 sink_events.lock().unwrap().push(event);
898 }));
899 let provider = Arc::new(CapturingProvider {
900 inner: MockProvider::with_results(script),
901 requests: Arc::clone(&requests),
902 });
903 let session = Session::new(provider, registry, vec![], config(), sink);
904 (session, requests, events)
905 }
906
907 fn user_text(message: &Message) -> Option<&str> {
908 match (message.role, message.content.as_slice()) {
909 (Role::User, [ContentBlock::Text { text }]) => Some(text.as_str()),
910 _ => None,
911 }
912 }
913
914 #[tokio::test]
915 async fn second_run_continues_the_conversation() {
916 let (mut s, requests, _e) = capturing_session_with(
917 vec![
918 Ok(text_turn("first answer")),
919 Ok(text_turn("second answer")),
920 ],
921 Registry::new(),
922 );
923 let r1 = s.run_text("q1").await;
924 let r2 = s.run_text("q2").await;
925 assert_eq!(r1.status, Status::Completed);
926 assert_eq!(r2.status, Status::Completed);
927 assert_eq!(r2.final_message.as_deref(), Some("second answer"));
928
929 let reqs = requests.lock().unwrap();
931 assert_eq!(reqs.len(), 2);
932 let run2 = &reqs[1];
933 assert_eq!(run2.len(), 3, "user q1, assistant, user q2: {run2:?}");
934 assert_eq!(user_text(&run2[0]), Some("q1"));
935 assert_eq!(run2[1].role, Role::Assistant);
936 assert_eq!(user_text(&run2[2]), Some("q2"));
937
938 let roles: Vec<Role> = s.history().iter().map(|m| m.role).collect();
940 assert_eq!(
941 roles,
942 vec![Role::User, Role::Assistant, Role::User, Role::Assistant]
943 );
944 }
945
946 #[tokio::test]
947 async fn init_emitted_once_across_runs_with_one_result_each() {
948 let (mut s, events) = session_with(
949 vec![Ok(text_turn("one")), Ok(text_turn("two"))],
950 Registry::new(),
951 config(),
952 );
953 let _ = s.run_text("q1").await;
954 let _ = s.run_text("q2").await;
955 let evs = dump(&events);
956 let inits = evs
957 .iter()
958 .filter(|e| matches!(e, Event::Init { .. }))
959 .count();
960 let results = evs
961 .iter()
962 .filter(|e| matches!(e, Event::Result { .. }))
963 .count();
964 assert_eq!(inits, 1, "Init is once per session, not per run");
965 assert_eq!(results, 2, "one Result per run");
966 assert!(
967 matches!(evs.first(), Some(Event::Init { .. })),
968 "Init still opens the stream"
969 );
970 }
971
972 #[tokio::test]
973 async fn report_counts_are_per_run_not_cumulative() {
974 let mut t1 = tool_turn("c1", "echo");
977 t1.usage = Usage {
978 input_tokens: 10,
979 output_tokens: 5,
980 ..Usage::default()
981 };
982 let t2 = text_turn("done one");
983 let mut t3 = text_turn("done two");
984 t3.usage = Usage {
985 input_tokens: 20,
986 output_tokens: 7,
987 ..Usage::default()
988 };
989 let (mut s, _e) = session_with(vec![Ok(t1), Ok(t2), Ok(t3)], echo_registry(), config());
990 let r1 = s.run_text("q1").await;
991 let r2 = s.run_text("q2").await;
992 assert_eq!(r1.turns, 2);
993 assert_eq!(r1.tool_calls.len(), 1);
994 assert_eq!(r2.turns, 1, "run 2 counts its own turns only");
995 assert!(r2.tool_calls.is_empty());
996 assert_eq!(r2.usage.input_tokens, 20, "usage is per-run");
997 assert_eq!(r2.usage.output_tokens, 7);
998 }
999
1000 #[tokio::test]
1003 async fn two_run_stream_reconstructs_the_full_conversation() {
1004 let (mut s, events) = session_with(
1005 vec![
1006 Ok(tool_turn("c1", "echo")),
1007 Ok(text_turn("done one")),
1008 Ok(text_turn("done two")),
1009 ],
1010 echo_registry(),
1011 config(),
1012 );
1013 let _ = s.run_text("q1").await;
1014 let _ = s.run_text("q2").await;
1015 let rebuilt: Conversation = reconstruct_conversation(&dump(&events));
1016 let roles: Vec<Role> = rebuilt.messages.iter().map(|m| m.role).collect();
1019 assert_eq!(
1020 roles,
1021 vec![
1022 Role::User,
1023 Role::Assistant,
1024 Role::User,
1025 Role::Assistant,
1026 Role::User,
1027 Role::Assistant,
1028 ]
1029 );
1030 assert_eq!(rebuilt.messages.as_slice(), s.history());
1032 }
1033
1034 #[tokio::test]
1037 async fn continues_after_model_error() {
1038 let (mut s, requests, _e) = capturing_session_with(
1039 vec![Err(ProviderError::ContextOverflow), Ok(text_turn("ok now"))],
1040 Registry::new(),
1041 );
1042 let r1 = s.run_text("q1").await;
1043 let r2 = s.run_text("q2").await;
1044 assert_eq!(r1.status, Status::ModelError);
1045 assert_eq!(r2.status, Status::Completed);
1046 let reqs = requests.lock().unwrap();
1048 let run2 = &reqs[1];
1049 assert_eq!(run2.len(), 2, "user q1 + user q2: {run2:?}");
1050 assert_eq!(user_text(&run2[0]), Some("q1"));
1051 assert_eq!(user_text(&run2[1]), Some("q2"));
1052 }
1053
1054 #[tokio::test]
1057 async fn continues_after_fatal_tool_error_with_valid_pairing() {
1058 let mut reg = Registry::new();
1059 reg.register("boom", Boom);
1060 let (mut s, requests, _e) = capturing_session_with(
1061 vec![Ok(tool_turn("c1", "boom")), Ok(text_turn("recovered"))],
1062 reg,
1063 );
1064 let r1 = s.run_text("q1").await;
1065 let r2 = s.run_text("q2").await;
1066 assert_eq!(r1.status, Status::Error);
1067 assert_eq!(r2.status, Status::Completed);
1068
1069 let reqs = requests.lock().unwrap();
1072 let run2 = &reqs[1];
1073 assert_eq!(run2.len(), 4, "q1, assistant, tool_result, q2: {run2:?}");
1074 assert!(
1075 run2[1]
1076 .content
1077 .iter()
1078 .any(|b| matches!(b, ContentBlock::ToolUse { id, .. } if id == "c1"))
1079 );
1080 assert!(run2[2].content.iter().any(|b| matches!(
1081 b,
1082 ContentBlock::ToolResult { tool_use_id, is_error: true, .. } if tool_use_id == "c1"
1083 )));
1084 assert_eq!(user_text(&run2[3]), Some("q2"));
1085 }
1086
1087 #[tokio::test]
1088 async fn usage_is_summed_across_turns() {
1089 let mut first = tool_turn("c1", "echo");
1090 first.usage = Usage {
1091 input_tokens: 10,
1092 output_tokens: 5,
1093 ..Usage::default()
1094 };
1095 let mut second = text_turn("done");
1096 second.usage = Usage {
1097 input_tokens: 20,
1098 output_tokens: 7,
1099 ..Usage::default()
1100 };
1101 let (mut s, _e) = session_with(vec![Ok(first), Ok(second)], echo_registry(), config());
1102 let report = s.run_text("go").await;
1103 assert_eq!(report.usage.input_tokens, 30);
1104 assert_eq!(report.usage.output_tokens, 12);
1105 }
1106}