1mod approve;
11mod config;
12mod instructions;
13mod run;
14mod session;
15mod sink;
16mod terminal;
17
18pub use approve::{AllowAll, ApprovalRequest, Approver, Decision};
19pub use config::EngineConfig;
20pub use session::Session;
21pub use sink::{EventSink, FnSink, NullSink};
22pub use tokio_util::sync::CancellationToken;
25
26#[cfg(test)]
27mod tests {
28 #![allow(clippy::unnecessary_literal_bound)]
31
32 use super::*;
33 use async_trait::async_trait;
34 use locode_protocol::{
35 ContentBlock, Conversation, Event, Message, ReasoningFormat, Role, Status, Usage,
36 reconstruct_conversation,
37 };
38 use locode_provider::{
39 Completion, ConversationRequest, MockProvider, Provider, ProviderError, StopReason,
40 };
41 use locode_tools::{Registry, Tool, ToolCtx, ToolError, ToolKind, ToolOutput};
42 use serde::Serialize;
43 use serde_json::{Value, json};
44 use std::sync::{Arc, Mutex};
45 use std::time::Duration;
46
47 #[derive(Serialize)]
50 struct EchoOut {
51 echoed: String,
52 }
53 impl ToolOutput for EchoOut {
54 fn to_prompt_text(&self) -> String {
55 self.echoed.clone()
56 }
57 }
58
59 struct Echo;
60 #[async_trait]
61 impl Tool for Echo {
62 type Args = Value;
63 type Output = EchoOut;
64 fn kind(&self) -> ToolKind {
65 ToolKind::Shell
66 }
67 fn description(&self) -> &str {
68 "echo"
69 }
70 async fn run(&self, _ctx: &ToolCtx, args: Value) -> Result<EchoOut, ToolError> {
71 Ok(EchoOut {
72 echoed: args.to_string(),
73 })
74 }
75 }
76
77 struct Boom;
78 #[async_trait]
79 impl Tool for Boom {
80 type Args = Value;
81 type Output = EchoOut;
82 fn kind(&self) -> ToolKind {
83 ToolKind::Shell
84 }
85 fn description(&self) -> &str {
86 "boom"
87 }
88 async fn run(&self, _ctx: &ToolCtx, _args: Value) -> Result<EchoOut, ToolError> {
89 Err(ToolError::Fatal("boom aborted the turn".into()))
90 }
91 }
92
93 fn text_turn(text: &str) -> Completion {
96 Completion {
97 content: vec![ContentBlock::Text { text: text.into() }],
98 usage: Usage::default(),
99 stop: StopReason::EndTurn,
100 }
101 }
102
103 fn tool_turn(id: &str, name: &str) -> Completion {
104 Completion {
105 content: vec![ContentBlock::ToolUse {
106 id: id.into(),
107 name: name.into(),
108 input: json!({}),
109 }],
110 usage: Usage::default(),
111 stop: StopReason::ToolUse,
112 }
113 }
114
115 fn config() -> EngineConfig {
116 EngineConfig {
117 session_id: "sess-1".into(),
118 harness: "grok".into(),
119 api_schema: "mock".into(),
120 model: "mock-1".into(),
121 max_turns: None,
122 resample_retries: 2,
123 resample_backoff: Duration::ZERO, instructions: locode_host::InstructionsConfig {
129 enabled: false,
130 ..Default::default()
131 },
132 ..EngineConfig::default()
133 }
134 }
135
136 fn session_with(
138 script: Vec<Result<Completion, ProviderError>>,
139 registry: Registry,
140 cfg: EngineConfig,
141 ) -> (Session, Arc<Mutex<Vec<Event>>>) {
142 let events = Arc::new(Mutex::new(Vec::new()));
143 let sink_events = Arc::clone(&events);
144 let sink = Box::new(FnSink(move |event| {
145 sink_events.lock().unwrap().push(event);
146 }));
147 let provider = Arc::new(MockProvider::with_results(script));
148 let session = Session::new(provider, registry, vec![], cfg, sink);
149 (session, events)
150 }
151
152 fn echo_registry() -> Registry {
153 let mut reg = Registry::new();
154 reg.register("echo", Echo);
155 reg
156 }
157
158 fn dump(events: &Arc<Mutex<Vec<Event>>>) -> Vec<Event> {
159 events.lock().unwrap().clone()
160 }
161
162 #[tokio::test]
165 async fn completed_with_no_tools() {
166 let (mut s, events) =
167 session_with(vec![Ok(text_turn("all done"))], Registry::new(), config());
168 let report = s.run_text("hi").await;
169 assert_eq!(report.status, Status::Completed);
170 assert_eq!(report.final_message.as_deref(), Some("all done"));
171 assert_eq!(report.turns, 1);
172 assert!(report.tool_calls.is_empty());
173 assert_eq!(report.api_schema, "mock");
174 let evs = dump(&events);
176 assert!(matches!(evs.first(), Some(Event::Init { .. })));
177 assert!(matches!(evs.last(), Some(Event::Result { .. })));
178 }
179
180 #[tokio::test]
183 async fn streaming_run_emits_text_deltas_and_the_whole_message() {
184 let mut cfg = config();
185 cfg.streaming = true;
186 let (mut s, events) = session_with(
187 vec![Ok(text_turn("hello streamed world"))],
188 Registry::new(),
189 cfg,
190 );
191 let report = s.run_text("hi").await;
192 assert_eq!(report.status, Status::Completed);
193 assert_eq!(
194 report.final_message.as_deref(),
195 Some("hello streamed world")
196 );
197
198 let evs = dump(&events);
199 let delta_text: String = evs
201 .iter()
202 .filter_map(|e| match e {
203 Event::MessageDelta { text } => Some(text.as_str()),
204 _ => None,
205 })
206 .collect();
207 assert_eq!(delta_text, "hello streamed world", "{evs:?}");
208 let n_deltas = evs
210 .iter()
211 .filter(|e| matches!(e, Event::MessageDelta { .. }))
212 .count();
213 assert!(n_deltas > 1, "expected multiple deltas, got {n_deltas}");
214 assert!(
216 evs.iter().any(|e| matches!(
217 e,
218 Event::Message { message } if message.role == Role::Assistant
219 )),
220 "whole assistant Message still emitted: {evs:?}"
221 );
222 let first_delta = evs
224 .iter()
225 .position(|e| matches!(e, Event::MessageDelta { .. }))
226 .expect("a delta");
227 let asst_msg = evs
228 .iter()
229 .position(
230 |e| matches!(e, Event::Message { message } if message.role == Role::Assistant),
231 )
232 .expect("assistant message");
233 assert!(
234 first_delta < asst_msg,
235 "deltas come before the whole message"
236 );
237 }
238
239 #[tokio::test]
240 async fn non_streaming_run_emits_no_deltas() {
241 let (mut s, events) =
242 session_with(vec![Ok(text_turn("no stream"))], Registry::new(), config());
243 let _ = s.run_text("hi").await;
244 let evs = dump(&events);
245 assert!(
246 !evs.iter().any(|e| matches!(e, Event::MessageDelta { .. })),
247 "default (non-streaming) run must not emit deltas: {evs:?}"
248 );
249 }
250
251 #[tokio::test]
252 async fn streaming_and_non_streaming_reports_match() {
253 let (mut a, _ea) = session_with(
254 vec![Ok(text_turn("same result"))],
255 Registry::new(),
256 config(),
257 );
258 let mut cfg = config();
259 cfg.streaming = true;
260 let (mut b, _eb) = session_with(vec![Ok(text_turn("same result"))], Registry::new(), cfg);
261 let ra = a.run_text("go").await;
262 let rb = b.run_text("go").await;
263 assert_eq!(ra.status, rb.status);
265 assert_eq!(ra.final_message, rb.final_message);
266 assert_eq!(ra.turns, rb.turns);
267 assert_eq!(ra.tool_calls.len(), rb.tool_calls.len());
268 }
269
270 #[tokio::test]
271 async fn tool_call_then_complete() {
272 let (mut s, _e) = session_with(
273 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
274 echo_registry(),
275 config(),
276 );
277 let report = s.run_text("go").await;
278 assert_eq!(report.status, Status::Completed);
279 assert_eq!(report.turns, 2);
280 assert_eq!(report.tool_calls.len(), 1);
281 assert!(report.tool_calls[0].ok);
282 assert_eq!(report.tool_calls[0].name, "echo");
283 }
284
285 #[tokio::test]
286 async fn hits_max_turns_after_dispatch() {
287 let mut cfg = config();
289 cfg.max_turns = Some(2);
290 let (mut s, _e) = session_with(
291 vec![
292 Ok(tool_turn("c1", "echo")),
293 Ok(tool_turn("c2", "echo")),
294 Ok(tool_turn("c3", "echo")),
295 ],
296 echo_registry(),
297 cfg,
298 );
299 let report = s.run_text("go").await;
300 assert_eq!(report.status, Status::MaxTurns);
301 assert_eq!(report.turns, 2);
302 assert_eq!(report.tool_calls.len(), 2);
303 }
304
305 #[tokio::test]
306 async fn model_error_after_bounded_retry() {
307 let script = vec![
309 Err(ProviderError::Transport("reset".into())),
310 Err(ProviderError::Transport("reset".into())),
311 Err(ProviderError::Transport("reset".into())),
312 ];
313 let (mut s, events) = session_with(script, Registry::new(), config());
314 let report = s.run_text("go").await;
315 assert_eq!(report.status, Status::ModelError);
316 assert!(report.error.is_some());
317 assert_eq!(report.turns, 0);
318 let retries = dump(&events)
320 .iter()
321 .filter(|e| matches!(e, Event::Error { .. }))
322 .count();
323 assert_eq!(retries, 2);
324 }
325
326 #[tokio::test]
327 async fn model_error_non_retryable_is_immediate() {
328 let (mut s, events) = session_with(
329 vec![Err(ProviderError::ContextOverflow)],
330 Registry::new(),
331 config(),
332 );
333 let report = s.run_text("go").await;
334 assert_eq!(report.status, Status::ModelError);
335 let retries = dump(&events)
336 .iter()
337 .filter(|e| matches!(e, Event::Error { .. }))
338 .count();
339 assert_eq!(retries, 0, "a non-retryable error must not resample");
340 }
341
342 #[tokio::test]
343 async fn fatal_tool_error_ends_the_run() {
344 let mut reg = Registry::new();
345 reg.register("boom", Boom);
346 let (mut s, _e) = session_with(vec![Ok(tool_turn("c1", "boom"))], reg, config());
347 let report = s.run_text("go").await;
348 assert_eq!(report.status, Status::Error);
349 assert!(report.error.is_some());
350 assert_eq!(report.tool_calls.len(), 1);
352 assert!(!report.tool_calls[0].ok);
353 }
354
355 #[tokio::test]
359 async fn empty_completion_resamples_then_succeeds() {
360 let empty = Completion {
361 content: vec![ContentBlock::Reasoning {
362 format: ReasoningFormat::Anthropic,
363 text: "thinking only".into(),
364 signature: Some("sig".into()),
365 payload: None,
366 }],
367 usage: Usage::default(),
368 stop: StopReason::MaxTokens,
369 };
370 let (mut session, _events) = session_with(
371 vec![Ok(empty), Ok(text_turn("recovered"))],
372 echo_registry(),
373 config(),
374 );
375 let report = session.run_text("go").await;
376 assert_eq!(report.status, Status::Completed);
377 assert_eq!(report.final_message.as_deref(), Some("recovered"));
378 assert_eq!(report.stop_reason.as_deref(), Some("end_turn"));
379 }
380
381 #[tokio::test]
382 async fn persistent_empty_completions_are_model_error() {
383 let empty = || Completion {
384 content: vec![],
385 usage: Usage::default(),
386 stop: StopReason::MaxTokens,
387 };
388 let (mut session, _events) = session_with(
390 vec![Ok(empty()), Ok(empty()), Ok(empty())],
391 echo_registry(),
392 config(),
393 );
394 let report = session.run_text("go").await;
395 assert_eq!(report.status, Status::ModelError);
396 assert!(
397 report
398 .error
399 .as_deref()
400 .unwrap_or("")
401 .contains("empty completion"),
402 "error names the cause: {:?}",
403 report.error
404 );
405 assert_eq!(report.stop_reason, None, "no completion was accepted");
406 }
407
408 #[tokio::test]
411 async fn mid_batch_abort_synthesizes_results() {
412 let mut reg = Registry::new();
415 reg.register("boom", Boom);
416 reg.register("echo", Echo);
417 let completion = Completion {
418 content: vec![
419 ContentBlock::ToolUse {
420 id: "c_boom".into(),
421 name: "boom".into(),
422 input: json!({}),
423 },
424 ContentBlock::ToolUse {
425 id: "c_echo".into(),
426 name: "echo".into(),
427 input: json!({}),
428 },
429 ],
430 usage: Usage::default(),
431 stop: StopReason::ToolUse,
432 };
433 let (mut s, events) = session_with(vec![Ok(completion)], reg, config());
434 let report = s.run_text("go").await;
435 assert_eq!(report.status, Status::Error);
436
437 let evs = dump(&events);
439 let answered: Vec<String> = evs
440 .iter()
441 .filter_map(|e| match e {
442 Event::Message { message } if message.role == Role::User => Some(&message.content),
443 _ => None,
444 })
445 .flatten()
446 .filter_map(|b| match b {
447 ContentBlock::ToolResult { tool_use_id, .. } => Some(tool_use_id.clone()),
448 _ => None,
449 })
450 .collect();
451 assert!(answered.iter().any(|id| id == "c_boom"));
452 assert!(
453 answered.iter().any(|id| id == "c_echo"),
454 "the un-run echo must be paired"
455 );
456 assert_eq!(report.tool_calls.len(), 1);
458 }
459
460 #[tokio::test]
463 async fn thinking_block_is_appended_verbatim() {
464 let completion = Completion {
465 content: vec![
466 ContentBlock::Reasoning {
467 format: ReasoningFormat::Anthropic,
468 text: "reasoning".into(),
469 signature: Some("sig-xyz".into()),
470 payload: None,
471 },
472 ContentBlock::Text {
473 text: "answer".into(),
474 },
475 ],
476 usage: Usage::default(),
477 stop: StopReason::EndTurn,
478 };
479 let (mut s, events) = session_with(vec![Ok(completion)], Registry::new(), config());
480 let report = s.run_text("think").await;
481 assert_eq!(report.status, Status::Completed);
482 assert_eq!(report.final_message.as_deref(), Some("answer"));
483 let has_thinking = dump(&events).iter().any(|e| match e {
485 Event::Message { message } if message.role == Role::Assistant => {
486 message.content.iter().any(|b| {
487 matches!(
488 b,
489 ContentBlock::Reasoning { signature: Some(sig), .. } if sig == "sig-xyz"
490 )
491 })
492 }
493 _ => false,
494 });
495 assert!(
496 has_thinking,
497 "thinking + signature must survive into history"
498 );
499 }
500
501 #[tokio::test]
502 async fn events_reconstruct_the_history() {
503 let (mut s, events) = session_with(
504 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
505 echo_registry(),
506 config(),
507 );
508 let _ = s.run_text("go").await;
509 let rebuilt: Conversation = reconstruct_conversation(&dump(&events));
510 let roles: Vec<Role> = rebuilt.messages.iter().map(|m| m.role).collect();
512 assert_eq!(
513 roles,
514 vec![Role::User, Role::Assistant, Role::User, Role::Assistant]
515 );
516 }
517
518 use std::sync::atomic::{AtomicUsize, Ordering};
521
522 struct Counting(Arc<AtomicUsize>);
524 #[async_trait]
525 impl Tool for Counting {
526 type Args = Value;
527 type Output = EchoOut;
528 fn kind(&self) -> ToolKind {
529 ToolKind::Shell
530 }
531 fn description(&self) -> &str {
532 "counting"
533 }
534 async fn run(&self, _ctx: &ToolCtx, _args: Value) -> Result<EchoOut, ToolError> {
535 self.0.fetch_add(1, Ordering::SeqCst);
536 Ok(EchoOut {
537 echoed: "ran".into(),
538 })
539 }
540 }
541
542 type SeenKinds = Arc<Mutex<Vec<(String, Option<ToolKind>)>>>;
543
544 struct DenyNamed {
547 deny: Vec<&'static str>,
548 seen_kinds: SeenKinds,
549 }
550 #[async_trait]
551 impl Approver for DenyNamed {
552 async fn decide(&self, request: &ApprovalRequest<'_>) -> Decision {
553 self.seen_kinds
554 .lock()
555 .unwrap()
556 .push((request.tool_name.to_owned(), request.kind));
557 if self.deny.contains(&request.tool_name) {
558 Decision::Deny {
559 reason: format!("{} is not allowed here", request.tool_name),
560 }
561 } else {
562 Decision::Allow
563 }
564 }
565 }
566
567 fn approvals(events: &Arc<Mutex<Vec<Event>>>) -> Vec<(String, String, String)> {
568 dump(events)
569 .iter()
570 .filter_map(|e| match e {
571 Event::Approval {
572 tool_use_id,
573 tool_name,
574 decision,
575 ..
576 } => Some((tool_use_id.clone(), tool_name.clone(), decision.clone())),
577 _ => None,
578 })
579 .collect()
580 }
581
582 #[tokio::test]
583 async fn deny_is_a_soft_paired_error_and_the_run_continues() {
584 let ran = Arc::new(AtomicUsize::new(0));
585 let mut reg = Registry::new();
586 reg.register("counting", Counting(Arc::clone(&ran)));
587 let (s, events) = session_with(
588 vec![Ok(tool_turn("c1", "counting")), Ok(text_turn("done"))],
589 reg,
590 config(),
591 );
592 let seen = Arc::new(Mutex::new(Vec::new()));
593 let mut s = s.with_approver(Arc::new(DenyNamed {
594 deny: vec!["counting"],
595 seen_kinds: Arc::clone(&seen),
596 }));
597 let report = s.run_text("go").await;
598
599 assert_eq!(report.status, Status::Completed);
601 assert_eq!(ran.load(Ordering::SeqCst), 0, "denied tool must not run");
602
603 assert_eq!(report.tool_calls.len(), 1);
605 let record = &report.tool_calls[0];
606 assert!(!record.ok);
607 assert_eq!(
608 record.denial_reason.as_deref(),
609 Some("counting is not allowed here")
610 );
611 assert_eq!(record.kind, "shell", "kind still recorded on denial");
612
613 let denied_result = dump(&events).iter().any(|e| match e {
615 Event::Message { message } => message.content.iter().any(|b| {
616 matches!(
617 b,
618 ContentBlock::ToolResult { tool_use_id, is_error: true, content, .. }
619 if tool_use_id == "c1"
620 && content.iter().any(|c| matches!(
621 c,
622 locode_protocol::ResultChunk::Text { text }
623 if text == "tool call denied: counting is not allowed here"
624 ))
625 )
626 }),
627 _ => false,
628 });
629 assert!(denied_result, "the model sees the denial reason, paired");
630
631 assert_eq!(
633 approvals(&events),
634 vec![("c1".into(), "counting".into(), "deny".into())]
635 );
636 }
637
638 #[tokio::test]
639 async fn deny_then_allow_within_one_batch_keeps_order_and_pairing() {
640 let ran = Arc::new(AtomicUsize::new(0));
641 let mut reg = Registry::new();
642 reg.register("blocked", Counting(Arc::clone(&ran)));
643 reg.register("echo", Echo);
644 let batch = Completion {
645 content: vec![
646 ContentBlock::ToolUse {
647 id: "c1".into(),
648 name: "blocked".into(),
649 input: json!({}),
650 },
651 ContentBlock::ToolUse {
652 id: "c2".into(),
653 name: "echo".into(),
654 input: json!({}),
655 },
656 ],
657 usage: Usage::default(),
658 stop: StopReason::ToolUse,
659 };
660 let (s, events) = session_with(vec![Ok(batch), Ok(text_turn("done"))], reg, config());
661 let mut s = s.with_approver(Arc::new(DenyNamed {
662 deny: vec!["blocked"],
663 seen_kinds: Arc::new(Mutex::new(Vec::new())),
664 }));
665 let report = s.run_text("go").await;
666 assert_eq!(report.status, Status::Completed);
667 assert_eq!(ran.load(Ordering::SeqCst), 0);
668
669 let pairs: Vec<(String, bool)> = dump(&events)
671 .iter()
672 .filter_map(|e| match e {
673 Event::Message { message } if message.role == Role::User => Some(&message.content),
674 _ => None,
675 })
676 .flatten()
677 .filter_map(|b| match b {
678 ContentBlock::ToolResult {
679 tool_use_id,
680 is_error,
681 ..
682 } => Some((tool_use_id.clone(), *is_error)),
683 _ => None,
684 })
685 .collect();
686 assert_eq!(pairs, vec![("c1".into(), true), ("c2".into(), false)]);
687
688 assert_eq!(report.tool_calls.len(), 2);
690 assert!(report.tool_calls[0].denial_reason.is_some());
691 assert_eq!(report.tool_calls[0].kind, "shell");
692 assert!(report.tool_calls[1].ok);
693 assert_eq!(report.tool_calls[1].denial_reason, None);
694
695 assert_eq!(
697 approvals(&events),
698 vec![
699 ("c1".into(), "blocked".into(), "deny".into()),
700 ("c2".into(), "echo".into(), "allow".into()),
701 ]
702 );
703 }
704
705 #[tokio::test]
706 async fn approval_request_carries_the_registry_kind() {
707 let seen = Arc::new(Mutex::new(Vec::new()));
708 let (s, _e) = session_with(
709 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
710 echo_registry(),
711 config(),
712 );
713 let mut s = s.with_approver(Arc::new(DenyNamed {
714 deny: vec![],
715 seen_kinds: Arc::clone(&seen),
716 }));
717 let _ = s.run_text("go").await;
718 let seen = seen.lock().unwrap();
719 assert_eq!(seen.len(), 1);
720 assert_eq!(seen[0].0, "echo");
721 assert_eq!(
722 seen[0].1,
723 Some(ToolKind::Shell),
724 "kind resolves from the registry pre-dispatch"
725 );
726 }
727
728 #[tokio::test]
732 async fn async_approver_suspends_the_call_until_resolved() {
733 struct OneshotApprover(Mutex<Option<tokio::sync::oneshot::Receiver<Decision>>>);
734 #[async_trait]
735 impl Approver for OneshotApprover {
736 async fn decide(&self, _request: &ApprovalRequest<'_>) -> Decision {
737 let rx = self.0.lock().unwrap().take().expect("one decision");
738 rx.await.expect("decider dropped")
739 }
740 }
741
742 let (tx, rx) = tokio::sync::oneshot::channel();
743 let (s, _e) = session_with(
744 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
745 echo_registry(),
746 config(),
747 );
748 let mut s = s.with_approver(Arc::new(OneshotApprover(Mutex::new(Some(rx)))));
749
750 let ui = tokio::spawn(async move {
752 tokio::task::yield_now().await;
753 let _ = tx.send(Decision::Allow);
754 });
755 let report = s.run_text("go").await;
756 ui.await.expect("ui task");
757 assert_eq!(report.status, Status::Completed);
758 assert_eq!(report.tool_calls.len(), 1);
759 assert!(report.tool_calls[0].ok);
760 }
761
762 #[tokio::test]
763 async fn allowed_calls_emit_approval_events_by_default() {
764 let (mut s, events) = session_with(
766 vec![Ok(tool_turn("c1", "echo")), Ok(text_turn("done"))],
767 echo_registry(),
768 config(),
769 );
770 let report = s.run_text("go").await;
771 assert_eq!(report.status, Status::Completed);
772 assert_eq!(
773 approvals(&events),
774 vec![("c1".into(), "echo".into(), "allow".into())]
775 );
776 assert_eq!(report.tool_calls[0].denial_reason, None);
778 }
779
780 struct HangingProvider;
785 #[async_trait]
786 impl Provider for HangingProvider {
787 #[allow(clippy::unnecessary_literal_bound)]
788 fn api_schema(&self) -> &str {
789 "mock"
790 }
791 async fn complete(
792 &self,
793 _request: &ConversationRequest,
794 ) -> Result<Completion, ProviderError> {
795 tokio::time::sleep(Duration::from_hours(1)).await;
796 Err(ProviderError::Transport("unreachable".into()))
797 }
798 }
799
800 struct WaitsForCancel;
803 #[async_trait]
804 impl Tool for WaitsForCancel {
805 type Args = Value;
806 type Output = EchoOut;
807 fn kind(&self) -> ToolKind {
808 ToolKind::Shell
809 }
810 fn description(&self) -> &str {
811 "waits"
812 }
813 async fn run(&self, ctx: &ToolCtx, _args: Value) -> Result<EchoOut, ToolError> {
814 ctx.cancel.cancelled().await;
815 Ok(EchoOut {
816 echoed: "stopped cooperatively".into(),
817 })
818 }
819 }
820
821 #[tokio::test]
822 async fn cancel_mid_sample_yields_cancelled_report() {
823 let events = Arc::new(Mutex::new(Vec::new()));
824 let sink_events = Arc::clone(&events);
825 let sink = Box::new(FnSink(move |event| {
826 sink_events.lock().unwrap().push(event);
827 }));
828 let mut s = Session::new(
829 Arc::new(HangingProvider),
830 Registry::new(),
831 vec![],
832 config(),
833 sink,
834 );
835 let handle = s.cancel_handle();
836 let canceller = tokio::spawn(async move {
837 tokio::time::sleep(Duration::from_millis(20)).await;
838 handle.cancel();
839 handle.cancel(); });
841 let report = s.run_text("go").await;
842 canceller.await.expect("canceller");
843
844 assert_eq!(report.status, Status::Cancelled);
845 assert_eq!(report.error, None, "cancelled is a stop, not a fault");
846 assert_eq!(report.final_message, None, "no assistant text this run");
847 assert_eq!(report.turns, 0, "no completion was accepted");
848 let roles: Vec<Role> = s.history().iter().map(|m| m.role).collect();
850 assert_eq!(roles, vec![Role::User]);
851 let evs = dump(&events);
853 assert!(
854 matches!(evs.last(), Some(Event::Result { report }) if report.status == Status::Cancelled)
855 );
856 }
857
858 #[tokio::test]
859 async fn cancel_mid_batch_pairs_the_rest_synthetically() {
860 let mut reg = Registry::new();
864 reg.register("waits", WaitsForCancel);
865 reg.register("echo", Echo);
866 let batch = Completion {
867 content: vec![
868 ContentBlock::ToolUse {
869 id: "c_wait".into(),
870 name: "waits".into(),
871 input: json!({}),
872 },
873 ContentBlock::ToolUse {
874 id: "c_echo".into(),
875 name: "echo".into(),
876 input: json!({}),
877 },
878 ],
879 usage: Usage::default(),
880 stop: StopReason::ToolUse,
881 };
882 let (s, events) = session_with(vec![Ok(batch)], reg, config());
883 let mut s = s; let handle = s.cancel_handle();
885 let canceller = tokio::spawn(async move {
886 tokio::time::sleep(Duration::from_millis(20)).await;
887 handle.cancel();
888 });
889 let report = s.run_text("go").await;
890 canceller.await.expect("canceller");
891
892 assert_eq!(report.status, Status::Cancelled);
893 assert_eq!(report.tool_calls.len(), 1);
896 assert_eq!(report.tool_calls[0].id, "c_wait");
897 assert!(report.tool_calls[0].ok);
898 assert_eq!(report.tool_calls[0].denial_reason, None);
899
900 let pairs: Vec<(String, bool)> = dump(&events)
902 .iter()
903 .filter_map(|e| match e {
904 Event::Message { message } if message.role == Role::User => Some(&message.content),
905 _ => None,
906 })
907 .flatten()
908 .filter_map(|b| match b {
909 ContentBlock::ToolResult {
910 tool_use_id,
911 is_error,
912 ..
913 } => Some((tool_use_id.clone(), *is_error)),
914 _ => None,
915 })
916 .collect();
917 assert_eq!(
918 pairs,
919 vec![("c_wait".into(), false), ("c_echo".into(), true)]
920 );
921 assert_eq!(
923 approvals(&events),
924 vec![("c_wait".into(), "waits".into(), "allow".into())]
925 );
926 }
927
928 #[tokio::test]
931 async fn cancelled_session_continues_on_the_next_run_with_a_fresh_token() {
932 let mut reg = Registry::new();
933 reg.register("waits", WaitsForCancel);
934 let (s, _e) = session_with(
935 vec![Ok(tool_turn("c1", "waits")), Ok(text_turn("second run"))],
936 reg,
937 config(),
938 );
939 let mut s = s;
940 let handle1 = s.cancel_handle();
941 let canceller = tokio::spawn(async move {
942 tokio::time::sleep(Duration::from_millis(20)).await;
943 handle1.cancel();
944 });
945 let r1 = s.run_text("q1").await;
946 canceller.await.expect("canceller");
947 assert_eq!(r1.status, Status::Cancelled);
948
949 assert!(!s.cancel_handle().is_cancelled());
952 let r2 = s.run_text("q2").await;
953 assert_eq!(r2.status, Status::Completed);
954 assert_eq!(r2.final_message.as_deref(), Some("second run"));
955 assert!(s.history().len() >= 4, "history: {:?}", s.history().len());
957 }
958
959 struct CapturingProvider {
964 inner: MockProvider,
965 requests: Arc<Mutex<Vec<Vec<Message>>>>,
966 }
967 #[async_trait]
968 impl Provider for CapturingProvider {
969 #[allow(clippy::unnecessary_literal_bound)]
970 fn api_schema(&self) -> &str {
971 "mock"
972 }
973 async fn complete(
974 &self,
975 request: &ConversationRequest,
976 ) -> Result<Completion, ProviderError> {
977 self.requests.lock().unwrap().push(request.messages.clone());
978 self.inner.complete(request).await
979 }
980 }
981
982 #[allow(clippy::type_complexity)]
984 fn capturing_session_with(
985 script: Vec<Result<Completion, ProviderError>>,
986 registry: Registry,
987 ) -> (
988 Session,
989 Arc<Mutex<Vec<Vec<Message>>>>,
990 Arc<Mutex<Vec<Event>>>,
991 ) {
992 let requests = Arc::new(Mutex::new(Vec::new()));
993 let events = Arc::new(Mutex::new(Vec::new()));
994 let sink_events = Arc::clone(&events);
995 let sink = Box::new(FnSink(move |event| {
996 sink_events.lock().unwrap().push(event);
997 }));
998 let provider = Arc::new(CapturingProvider {
999 inner: MockProvider::with_results(script),
1000 requests: Arc::clone(&requests),
1001 });
1002 let session = Session::new(provider, registry, vec![], config(), sink);
1003 (session, requests, events)
1004 }
1005
1006 fn user_text(message: &Message) -> Option<&str> {
1007 match (message.role, message.content.as_slice()) {
1008 (Role::User, [ContentBlock::Text { text }]) => Some(text.as_str()),
1009 _ => None,
1010 }
1011 }
1012
1013 #[tokio::test]
1014 async fn second_run_continues_the_conversation() {
1015 let (mut s, requests, _e) = capturing_session_with(
1016 vec![
1017 Ok(text_turn("first answer")),
1018 Ok(text_turn("second answer")),
1019 ],
1020 Registry::new(),
1021 );
1022 let r1 = s.run_text("q1").await;
1023 let r2 = s.run_text("q2").await;
1024 assert_eq!(r1.status, Status::Completed);
1025 assert_eq!(r2.status, Status::Completed);
1026 assert_eq!(r2.final_message.as_deref(), Some("second answer"));
1027
1028 let reqs = requests.lock().unwrap();
1030 assert_eq!(reqs.len(), 2);
1031 let run2 = &reqs[1];
1032 assert_eq!(run2.len(), 3, "user q1, assistant, user q2: {run2:?}");
1033 assert_eq!(user_text(&run2[0]), Some("q1"));
1034 assert_eq!(run2[1].role, Role::Assistant);
1035 assert_eq!(user_text(&run2[2]), Some("q2"));
1036
1037 let roles: Vec<Role> = s.history().iter().map(|m| m.role).collect();
1039 assert_eq!(
1040 roles,
1041 vec![Role::User, Role::Assistant, Role::User, Role::Assistant]
1042 );
1043 }
1044
1045 fn capturing_with_cfg(
1049 script: Vec<Result<Completion, ProviderError>>,
1050 cfg: EngineConfig,
1051 ) -> (Session, Arc<Mutex<Vec<Vec<Message>>>>) {
1052 let requests = Arc::new(Mutex::new(Vec::new()));
1053 let provider = Arc::new(CapturingProvider {
1054 inner: MockProvider::with_results(script),
1055 requests: Arc::clone(&requests),
1056 });
1057 let session = Session::new(provider, Registry::new(), vec![], cfg, Box::new(NullSink));
1058 (session, requests)
1059 }
1060
1061 fn instr_config(cwd: std::path::PathBuf) -> EngineConfig {
1064 EngineConfig {
1065 cwd,
1066 instructions: locode_host::InstructionsConfig {
1067 global_file: false,
1068 ..Default::default()
1069 },
1070 ..config()
1071 }
1072 }
1073
1074 fn reminder_text(msgs: &[Message]) -> Option<String> {
1076 msgs.iter()
1077 .find_map(|m| match (m.role, m.content.as_slice()) {
1078 (Role::User, [ContentBlock::Text { text }])
1079 if text.starts_with("<system-reminder>") =>
1080 {
1081 Some(text.clone())
1082 }
1083 _ => None,
1084 })
1085 }
1086
1087 fn reminder_count(msgs: &[Message]) -> usize {
1088 msgs.iter()
1089 .filter(|m| {
1090 matches!(
1091 (m.role, m.content.as_slice()),
1092 (Role::User, [ContentBlock::Text { text }]) if text.starts_with("<system-reminder>")
1093 )
1094 })
1095 .count()
1096 }
1097
1098 #[tokio::test]
1099 async fn project_instructions_injected_once_before_prompt() {
1100 let dir = tempfile::tempdir().unwrap();
1101 let root = std::fs::canonicalize(dir.path()).unwrap();
1102 std::fs::create_dir(root.join(".git")).unwrap();
1103 std::fs::write(root.join("AGENTS.md"), "be terse").unwrap();
1104
1105 let (mut s, requests) = capturing_with_cfg(
1106 vec![Ok(text_turn("ok1")), Ok(text_turn("ok2"))],
1107 instr_config(root),
1108 );
1109 s.run_text("q1").await;
1110 s.run_text("q2").await;
1111
1112 let reqs = requests.lock().unwrap();
1113 let run1 = &reqs[0];
1115 let rem = reminder_text(run1).expect("instructions injected on run 1");
1116 assert!(rem.contains("## From:"), "labeled: {rem}");
1117 assert!(rem.contains("be terse"), "content present: {rem}");
1118 let rem_idx = run1
1119 .iter()
1120 .position(|m| reminder_text(std::slice::from_ref(m)).is_some());
1121 let q1_idx = run1.iter().position(|m| user_text(m) == Some("q1"));
1122 assert!(rem_idx < q1_idx, "reminder comes before the prompt");
1123
1124 assert_eq!(reminder_count(&reqs[1]), 1, "not re-injected on run 2");
1126 }
1127
1128 #[tokio::test]
1129 async fn project_instructions_absent_when_disabled() {
1130 let dir = tempfile::tempdir().unwrap();
1131 let root = std::fs::canonicalize(dir.path()).unwrap();
1132 std::fs::create_dir(root.join(".git")).unwrap();
1133 std::fs::write(root.join("AGENTS.md"), "be terse").unwrap();
1134 let mut cfg = instr_config(root);
1135 cfg.instructions.enabled = false;
1136
1137 let (mut s, requests) = capturing_with_cfg(vec![Ok(text_turn("ok"))], cfg);
1138 s.run_text("q").await;
1139 assert!(reminder_text(&requests.lock().unwrap()[0]).is_none());
1140 }
1141
1142 #[tokio::test]
1143 async fn project_instructions_absent_when_no_agents_md() {
1144 let dir = tempfile::tempdir().unwrap();
1145 let root = std::fs::canonicalize(dir.path()).unwrap();
1146 let (mut s, requests) = capturing_with_cfg(vec![Ok(text_turn("ok"))], instr_config(root));
1148 s.run_text("q").await;
1149 assert!(reminder_text(&requests.lock().unwrap()[0]).is_none());
1150 }
1151
1152 #[tokio::test]
1153 async fn project_instructions_replace_banner_on_edit() {
1154 let dir = tempfile::tempdir().unwrap();
1155 let root = std::fs::canonicalize(dir.path()).unwrap();
1156 std::fs::create_dir(root.join(".git")).unwrap();
1157 let agents = root.join("AGENTS.md");
1158 std::fs::write(&agents, "v1 rules").unwrap();
1159
1160 let (mut s, requests) = capturing_with_cfg(
1161 vec![Ok(text_turn("ok1")), Ok(text_turn("ok2"))],
1162 instr_config(root),
1163 );
1164 s.run_text("q1").await;
1165 std::fs::write(&agents, "v2 rules").unwrap(); s.run_text("q2").await;
1167
1168 let reqs = requests.lock().unwrap();
1169 let run2 = &reqs[1];
1170 let banner = run2
1171 .iter()
1172 .find_map(|m| match (m.role, m.content.as_slice()) {
1173 (Role::User, [ContentBlock::Text { text }])
1174 if text.contains("replace all previously provided") =>
1175 {
1176 Some(text.clone())
1177 }
1178 _ => None,
1179 })
1180 .expect("replace banner on edit");
1181 assert!(banner.contains("v2 rules"), "new content: {banner}");
1182 assert!(!banner.contains("v1 rules"), "not the old content");
1183 assert_eq!(reminder_count(run2), 2);
1185 }
1186
1187 #[tokio::test]
1188 async fn project_instructions_removal_banner_on_delete() {
1189 let dir = tempfile::tempdir().unwrap();
1190 let root = std::fs::canonicalize(dir.path()).unwrap();
1191 std::fs::create_dir(root.join(".git")).unwrap();
1192 let agents = root.join("AGENTS.md");
1193 std::fs::write(&agents, "rules").unwrap();
1194
1195 let (mut s, requests) = capturing_with_cfg(
1196 vec![Ok(text_turn("ok1")), Ok(text_turn("ok2"))],
1197 instr_config(root),
1198 );
1199 s.run_text("q1").await;
1200 std::fs::remove_file(&agents).unwrap(); s.run_text("q2").await;
1202
1203 let reqs = requests.lock().unwrap();
1204 assert!(
1205 reqs[1].iter().any(|m| matches!(
1206 (m.role, m.content.as_slice()),
1207 (Role::User, [ContentBlock::Text { text }]) if text.contains("no longer apply")
1208 )),
1209 "removal notice on delete"
1210 );
1211 }
1212
1213 #[tokio::test]
1214 async fn project_instructions_not_reinjected_when_unchanged() {
1215 let dir = tempfile::tempdir().unwrap();
1216 let root = std::fs::canonicalize(dir.path()).unwrap();
1217 std::fs::create_dir(root.join(".git")).unwrap();
1218 std::fs::write(root.join("AGENTS.md"), "stable").unwrap();
1219
1220 let (mut s, requests) = capturing_with_cfg(
1221 vec![
1222 Ok(text_turn("ok1")),
1223 Ok(text_turn("ok2")),
1224 Ok(text_turn("ok3")),
1225 ],
1226 instr_config(root),
1227 );
1228 s.run_text("q1").await;
1229 s.run_text("q2").await;
1230 s.run_text("q3").await;
1231 assert_eq!(reminder_count(&requests.lock().unwrap()[2]), 1);
1233 }
1234
1235 #[tokio::test]
1236 async fn init_emitted_once_across_runs_with_one_result_each() {
1237 let (mut s, events) = session_with(
1238 vec![Ok(text_turn("one")), Ok(text_turn("two"))],
1239 Registry::new(),
1240 config(),
1241 );
1242 let _ = s.run_text("q1").await;
1243 let _ = s.run_text("q2").await;
1244 let evs = dump(&events);
1245 let inits = evs
1246 .iter()
1247 .filter(|e| matches!(e, Event::Init { .. }))
1248 .count();
1249 let results = evs
1250 .iter()
1251 .filter(|e| matches!(e, Event::Result { .. }))
1252 .count();
1253 assert_eq!(inits, 1, "Init is once per session, not per run");
1254 assert_eq!(results, 2, "one Result per run");
1255 assert!(
1256 matches!(evs.first(), Some(Event::Init { .. })),
1257 "Init still opens the stream"
1258 );
1259 }
1260
1261 #[tokio::test]
1262 async fn report_counts_are_per_run_not_cumulative() {
1263 let mut t1 = tool_turn("c1", "echo");
1266 t1.usage = Usage {
1267 input_tokens: 10,
1268 output_tokens: 5,
1269 ..Usage::default()
1270 };
1271 let t2 = text_turn("done one");
1272 let mut t3 = text_turn("done two");
1273 t3.usage = Usage {
1274 input_tokens: 20,
1275 output_tokens: 7,
1276 ..Usage::default()
1277 };
1278 let (mut s, _e) = session_with(vec![Ok(t1), Ok(t2), Ok(t3)], echo_registry(), config());
1279 let r1 = s.run_text("q1").await;
1280 let r2 = s.run_text("q2").await;
1281 assert_eq!(r1.turns, 2);
1282 assert_eq!(r1.tool_calls.len(), 1);
1283 assert_eq!(r2.turns, 1, "run 2 counts its own turns only");
1284 assert!(r2.tool_calls.is_empty());
1285 assert_eq!(r2.usage.input_tokens, 20, "usage is per-run");
1286 assert_eq!(r2.usage.output_tokens, 7);
1287 }
1288
1289 #[tokio::test]
1292 async fn two_run_stream_reconstructs_the_full_conversation() {
1293 let (mut s, events) = session_with(
1294 vec![
1295 Ok(tool_turn("c1", "echo")),
1296 Ok(text_turn("done one")),
1297 Ok(text_turn("done two")),
1298 ],
1299 echo_registry(),
1300 config(),
1301 );
1302 let _ = s.run_text("q1").await;
1303 let _ = s.run_text("q2").await;
1304 let rebuilt: Conversation = reconstruct_conversation(&dump(&events));
1305 let roles: Vec<Role> = rebuilt.messages.iter().map(|m| m.role).collect();
1308 assert_eq!(
1309 roles,
1310 vec![
1311 Role::User,
1312 Role::Assistant,
1313 Role::User,
1314 Role::Assistant,
1315 Role::User,
1316 Role::Assistant,
1317 ]
1318 );
1319 assert_eq!(rebuilt.messages.as_slice(), s.history());
1321 }
1322
1323 #[tokio::test]
1326 async fn continues_after_model_error() {
1327 let (mut s, requests, _e) = capturing_session_with(
1328 vec![Err(ProviderError::ContextOverflow), Ok(text_turn("ok now"))],
1329 Registry::new(),
1330 );
1331 let r1 = s.run_text("q1").await;
1332 let r2 = s.run_text("q2").await;
1333 assert_eq!(r1.status, Status::ModelError);
1334 assert_eq!(r2.status, Status::Completed);
1335 let reqs = requests.lock().unwrap();
1337 let run2 = &reqs[1];
1338 assert_eq!(run2.len(), 2, "user q1 + user q2: {run2:?}");
1339 assert_eq!(user_text(&run2[0]), Some("q1"));
1340 assert_eq!(user_text(&run2[1]), Some("q2"));
1341 }
1342
1343 #[tokio::test]
1346 async fn continues_after_fatal_tool_error_with_valid_pairing() {
1347 let mut reg = Registry::new();
1348 reg.register("boom", Boom);
1349 let (mut s, requests, _e) = capturing_session_with(
1350 vec![Ok(tool_turn("c1", "boom")), Ok(text_turn("recovered"))],
1351 reg,
1352 );
1353 let r1 = s.run_text("q1").await;
1354 let r2 = s.run_text("q2").await;
1355 assert_eq!(r1.status, Status::Error);
1356 assert_eq!(r2.status, Status::Completed);
1357
1358 let reqs = requests.lock().unwrap();
1361 let run2 = &reqs[1];
1362 assert_eq!(run2.len(), 4, "q1, assistant, tool_result, q2: {run2:?}");
1363 assert!(
1364 run2[1]
1365 .content
1366 .iter()
1367 .any(|b| matches!(b, ContentBlock::ToolUse { id, .. } if id == "c1"))
1368 );
1369 assert!(run2[2].content.iter().any(|b| matches!(
1370 b,
1371 ContentBlock::ToolResult { tool_use_id, is_error: true, .. } if tool_use_id == "c1"
1372 )));
1373 assert_eq!(user_text(&run2[3]), Some("q2"));
1374 }
1375
1376 #[tokio::test]
1377 async fn usage_is_summed_across_turns() {
1378 let mut first = tool_turn("c1", "echo");
1379 first.usage = Usage {
1380 input_tokens: 10,
1381 output_tokens: 5,
1382 ..Usage::default()
1383 };
1384 let mut second = text_turn("done");
1385 second.usage = Usage {
1386 input_tokens: 20,
1387 output_tokens: 7,
1388 ..Usage::default()
1389 };
1390 let (mut s, _e) = session_with(vec![Ok(first), Ok(second)], echo_registry(), config());
1391 let report = s.run_text("go").await;
1392 assert_eq!(report.usage.input_tokens, 30);
1393 assert_eq!(report.usage.output_tokens, 12);
1394 }
1395}