1use std::{
4 collections::{HashMap, VecDeque},
5 fmt,
6 future::Future,
7 path::PathBuf,
8 sync::{Arc, Mutex, PoisonError},
9 time::Duration,
10};
11
12use async_trait::async_trait;
13use serde::{Deserialize, Serialize};
14use serde_json::Value;
15use thiserror::Error;
16use tokio_util::sync::CancellationToken;
17
18#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
19#[serde(tag = "role", rename_all = "snake_case")]
20pub enum Message {
21 User {
22 content: String,
23 },
24 Assistant {
25 content: String,
26 #[serde(default, skip_serializing_if = "Vec::is_empty")]
27 tool_calls: Vec<ToolCall>,
28 },
29 Tool {
30 call_id: String,
31 name: String,
32 content: String,
33 is_error: bool,
34 },
35 HistoryNote {
36 content: String,
37 },
38}
39
40impl Message {
41 fn estimated_tokens(&self, bytes_per_token: usize) -> usize {
42 let bytes = serde_json::to_vec(self).map_or(0, |value| value.len());
43 bytes.div_ceil(bytes_per_token).saturating_add(4)
44 }
45}
46
47#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
48pub struct ToolCall {
49 pub id: String,
50 pub name: String,
51 pub arguments: Value,
52}
53
54#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
55pub struct ToolSpec {
56 pub name: String,
57 pub description: String,
58 pub parameters: Value,
59}
60
61#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
62#[serde(rename_all = "snake_case")]
63pub enum ToolRisk {
64 ReadOnly,
65 Filesystem,
66 Process,
67 Delegate,
68 Network,
71}
72
73impl ToolRisk {
74 pub fn as_str(self) -> &'static str {
75 match self {
76 Self::ReadOnly => "read_only",
77 Self::Filesystem => "filesystem",
78 Self::Process => "process",
79 Self::Delegate => "delegate",
80 Self::Network => "network",
81 }
82 }
83}
84
85#[derive(Debug, Clone)]
86pub struct ToolContext {
87 pub workspace: PathBuf,
88 pub cancellation: CancellationToken,
89 pub progress: ProgressSink,
91}
92
93impl ToolContext {
94 pub fn new(workspace: PathBuf, cancellation: CancellationToken) -> Self {
96 Self {
97 workspace,
98 cancellation,
99 progress: ProgressSink::default(),
100 }
101 }
102}
103
104pub const MAX_PROGRESS_LINE_BYTES: usize = 200;
106pub const MAX_PROGRESS_EVENT_BYTES: usize = 512;
109pub const PROGRESS_INTERVAL: Duration = Duration::from_millis(500);
111
112#[derive(Clone, Default)]
118pub struct ProgressSink {
119 pending: Option<Arc<Mutex<PendingProgress>>>,
120}
121
122impl fmt::Debug for ProgressSink {
123 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
124 formatter
125 .debug_struct("ProgressSink")
126 .field("enabled", &self.is_enabled())
127 .finish()
128 }
129}
130
131impl ProgressSink {
132 pub fn buffered() -> Self {
134 Self {
135 pending: Some(Arc::default()),
136 }
137 }
138
139 pub fn is_enabled(&self) -> bool {
140 self.pending.is_some()
141 }
142
143 pub fn report(&self, text: &str) {
146 let Some(pending) = &self.pending else {
147 return;
148 };
149 let line = progress_line(text);
150 if !line.is_empty() {
151 pending
152 .lock()
153 .unwrap_or_else(PoisonError::into_inner)
154 .push(line);
155 }
156 }
157
158 pub fn take(&self) -> Option<String> {
161 self.pending
162 .as_ref()?
163 .lock()
164 .unwrap_or_else(PoisonError::into_inner)
165 .take()
166 }
167}
168
169const PROGRESS_ELIDED: &str = "…";
171
172#[derive(Debug, Default)]
173struct PendingProgress {
174 lines: VecDeque<String>,
175 bytes: usize,
177 dropped: bool,
178}
179
180impl PendingProgress {
181 fn push(&mut self, line: String) {
182 self.bytes += line.len() + usize::from(!self.lines.is_empty());
183 self.lines.push_back(line);
184 let budget = MAX_PROGRESS_EVENT_BYTES - PROGRESS_ELIDED.len() - 1;
186 while self.bytes > budget && self.lines.len() > 1 {
187 if let Some(oldest) = self.lines.pop_front() {
188 self.bytes -= oldest.len() + 1;
189 self.dropped = true;
190 }
191 }
192 }
193
194 fn take(&mut self) -> Option<String> {
195 if self.lines.is_empty() {
196 return None;
197 }
198 let mut text = String::with_capacity(self.bytes + PROGRESS_ELIDED.len() + 1);
199 if std::mem::take(&mut self.dropped) {
200 text.push_str(PROGRESS_ELIDED);
201 text.push('\n');
202 }
203 for (index, line) in self.lines.drain(..).enumerate() {
204 if index > 0 {
205 text.push('\n');
206 }
207 text.push_str(&line);
208 }
209 self.bytes = 0;
210 Some(text)
211 }
212}
213
214fn progress_line(text: &str) -> String {
217 let mut line = String::new();
218 for word in text
219 .split(|character: char| character.is_whitespace() || character.is_control())
220 .filter(|word| !word.is_empty())
221 {
222 if !line.is_empty() {
223 line.push(' ');
224 }
225 line.push_str(word);
226 if line.len() > MAX_PROGRESS_LINE_BYTES {
227 break;
228 }
229 }
230 if line.len() <= MAX_PROGRESS_LINE_BYTES {
231 return line;
232 }
233 let mut end = MAX_PROGRESS_LINE_BYTES - PROGRESS_ELIDED.len();
234 while !line.is_char_boundary(end) {
235 end -= 1;
236 }
237 line.truncate(end);
238 line.push_str(PROGRESS_ELIDED);
239 line
240}
241
242#[derive(Debug, Clone, PartialEq, Eq)]
243pub struct ToolOutput {
244 pub content: String,
245 pub is_error: bool,
246 pub truncated: bool,
247}
248
249impl ToolOutput {
250 pub fn success(content: impl Into<String>) -> Self {
251 Self {
252 content: content.into(),
253 is_error: false,
254 truncated: false,
255 }
256 }
257
258 pub fn failure(content: impl Into<String>) -> Self {
259 Self {
260 content: content.into(),
261 is_error: true,
262 truncated: false,
263 }
264 }
265}
266
267#[derive(Debug, Error)]
268#[error("{0}")]
269pub struct ToolError(pub String);
270
271#[async_trait]
272pub trait Tool: Send + Sync {
273 fn spec(&self) -> ToolSpec;
274 fn risk(&self, arguments: &Value) -> Result<ToolRisk, ToolError>;
275 fn approval_summary(&self, arguments: &Value) -> Result<String, ToolError>;
276 async fn execute(
277 &self,
278 arguments: Value,
279 context: ToolContext,
280 ) -> Result<ToolOutput, ToolError>;
281}
282
283#[derive(Default)]
284pub struct ToolRegistry {
285 tools: HashMap<String, Arc<dyn Tool>>,
286}
287
288impl ToolRegistry {
289 pub fn register(&mut self, tool: Arc<dyn Tool>) -> Result<(), ToolError> {
290 let name = tool.spec().name;
291 if self.tools.contains_key(&name) {
292 return Err(ToolError(format!("duplicate tool name: {name}")));
293 }
294 self.tools.insert(name, tool);
295 Ok(())
296 }
297
298 pub fn get(&self, name: &str) -> Option<Arc<dyn Tool>> {
299 self.tools.get(name).cloned()
300 }
301
302 pub fn specs(&self) -> Vec<ToolSpec> {
303 let mut specs: Vec<_> = self.tools.values().map(|tool| tool.spec()).collect();
304 specs.sort_by(|a, b| a.name.cmp(&b.name));
305 specs
306 }
307}
308
309#[derive(Debug, Clone, Default, PartialEq, Eq)]
310pub struct Usage {
311 pub input_tokens: Option<u64>,
312 pub output_tokens: Option<u64>,
313}
314
315impl Usage {
316 fn add(&mut self, other: &Self) {
317 self.input_tokens = add_optional(self.input_tokens, other.input_tokens);
318 self.output_tokens = add_optional(self.output_tokens, other.output_tokens);
319 }
320}
321
322fn add_optional(left: Option<u64>, right: Option<u64>) -> Option<u64> {
323 match (left, right) {
324 (None, None) => None,
325 (left, right) => Some(left.unwrap_or(0).saturating_add(right.unwrap_or(0))),
326 }
327}
328
329#[derive(Debug, Clone)]
330pub struct ProviderRequest {
331 pub system_prompt: String,
332 pub messages: Vec<Message>,
333 pub tools: Vec<ToolSpec>,
334}
335
336#[derive(Debug, Clone)]
337pub struct AssistantResponse {
338 pub content: String,
339 pub tool_calls: Vec<ToolCall>,
340 pub usage: Usage,
341}
342
343#[derive(Debug, Clone, Copy, PartialEq, Eq)]
344pub enum ProviderErrorKind {
345 Provider,
346 ResponseLimit,
347 ToolLimit,
348 Cancelled,
349}
350
351#[derive(Debug, Error)]
352#[error("{message}")]
353pub struct ProviderError {
354 pub kind: ProviderErrorKind,
355 pub message: String,
356}
357
358impl ProviderError {
359 pub fn new(kind: ProviderErrorKind, message: impl Into<String>) -> Self {
360 Self {
361 kind,
362 message: message.into(),
363 }
364 }
365}
366
367#[async_trait]
368pub trait TextDeltaSink: Send + Sync {
369 async fn push(&self, delta: &str) -> Result<(), ProviderError>;
370}
371
372#[async_trait]
373pub trait Provider: Send + Sync {
374 fn model(&self) -> &str;
375
376 async fn complete(
377 &self,
378 request: ProviderRequest,
379 deltas: Arc<dyn TextDeltaSink>,
380 cancellation: CancellationToken,
381 ) -> Result<AssistantResponse, ProviderError>;
382}
383
384#[derive(Debug, Clone)]
385pub enum CoreEvent {
386 AssistantDelta {
387 content: String,
388 },
389 AssistantCompleted {
390 content: String,
391 },
392 ToolProposed {
393 call_id: String,
394 name: String,
395 arguments: Value,
396 },
397 ToolStarted {
398 call_id: String,
399 name: String,
400 },
401 ToolProgress {
403 call_id: String,
404 text: String,
405 },
406 ToolCompleted {
407 call_id: String,
408 name: String,
409 output: ToolOutput,
410 },
411 ContextCompacted {
412 before_tokens: usize,
413 after_tokens: usize,
414 removed_messages: usize,
415 },
416 SessionTrimmed {
417 removed_messages: usize,
418 history_bytes: usize,
419 },
420}
421
422#[async_trait]
423pub trait EventSink: Send + Sync {
424 async fn emit(&self, event: CoreEvent) -> Result<(), AgentError>;
425}
426
427#[derive(Debug, Clone)]
428pub struct ApprovalRequest {
429 pub call_id: String,
430 pub name: String,
431 pub risk: ToolRisk,
432 pub cwd: PathBuf,
433 pub summary: String,
434}
435
436#[async_trait]
437pub trait ApprovalGate: Send + Sync {
438 async fn approve(
439 &self,
440 request: ApprovalRequest,
441 cancellation: CancellationToken,
442 ) -> Result<bool, AgentError>;
443}
444
445#[derive(Debug, Clone)]
446pub struct ContextConfig {
447 pub max_tokens: usize,
448 pub reserve_output_tokens: usize,
449 pub safety_margin_tokens: usize,
450 pub bytes_per_token: usize,
451 pub summary_max_chars: usize,
452}
453
454impl Default for ContextConfig {
455 fn default() -> Self {
456 Self {
457 max_tokens: 128_000,
458 reserve_output_tokens: 8_192,
459 safety_margin_tokens: 2_048,
460 bytes_per_token: 3,
461 summary_max_chars: 6_000,
462 }
463 }
464}
465
466#[derive(Debug, Clone)]
467pub struct ContextSelection {
468 pub messages: Vec<Message>,
469 pub before_tokens: usize,
470 pub after_tokens: usize,
471 pub removed_messages: usize,
472}
473
474#[derive(Debug, Error)]
475#[error("{0}")]
476pub struct ContextError(pub String);
477
478pub trait ContextPolicy: Send + Sync {
479 fn select(
480 &self,
481 history: &[Message],
482 system_prompt: &str,
483 tools: &[ToolSpec],
484 ) -> Result<ContextSelection, ContextError>;
485}
486
487pub struct BudgetContextPolicy {
488 config: ContextConfig,
489}
490
491impl BudgetContextPolicy {
492 pub fn new(config: ContextConfig) -> Result<Self, ContextError> {
493 if config.bytes_per_token == 0 {
494 return Err(ContextError(
495 "context.bytes_per_token must be positive".into(),
496 ));
497 }
498 if config
499 .reserve_output_tokens
500 .saturating_add(config.safety_margin_tokens)
501 >= config.max_tokens
502 {
503 return Err(ContextError(
504 "context reserve and safety margin consume the model window".into(),
505 ));
506 }
507 Ok(Self { config })
508 }
509
510 fn string_tokens(&self, value: &str) -> usize {
511 value.len().div_ceil(self.config.bytes_per_token)
512 }
513
514 fn group_messages(history: &[Message]) -> Vec<Vec<Message>> {
515 let mut groups: Vec<Vec<Message>> = Vec::new();
516 for message in history {
517 if matches!(message, Message::User { .. }) || groups.is_empty() {
518 groups.push(Vec::new());
519 }
520 groups
521 .last_mut()
522 .expect("a group was just created")
523 .push(message.clone());
524 }
525 groups
526 }
527
528 fn summarize(&self, messages: &[Message]) -> String {
529 let mut output = format!(
530 "[SCV compacted {} earlier messages. Bounded extracts follow.]\n",
531 messages.len()
532 );
533 for message in messages {
534 let (label, content) = match message {
535 Message::User { content } => ("user", content.as_str()),
536 Message::Assistant { content, .. } => ("assistant", content.as_str()),
537 Message::Tool {
538 name,
539 content,
540 is_error,
541 ..
542 } => {
543 let status = if *is_error { "failed" } else { "ok" };
544 output.push_str(&format!("tool {name} ({status}): "));
545 ("", content.as_str())
546 }
547 Message::HistoryNote { content } => ("earlier", content.as_str()),
548 };
549 if !label.is_empty() {
550 output.push_str(label);
551 output.push_str(": ");
552 }
553 let tail = char_tail(content, 240);
554 output.push_str(&tail.replace('\n', " "));
555 output.push('\n');
556 if output.chars().count() >= self.config.summary_max_chars {
557 break;
558 }
559 }
560 truncate_chars(&output, self.config.summary_max_chars)
561 }
562}
563
564impl ContextPolicy for BudgetContextPolicy {
565 fn select(
566 &self,
567 history: &[Message],
568 system_prompt: &str,
569 tools: &[ToolSpec],
570 ) -> Result<ContextSelection, ContextError> {
571 if history.is_empty() {
572 return Ok(ContextSelection {
573 messages: Vec::new(),
574 before_tokens: 0,
575 after_tokens: 0,
576 removed_messages: 0,
577 });
578 }
579 let tools_bytes = serde_json::to_vec(tools).map_or(0, |value| value.len());
580 let static_tokens = self
581 .string_tokens(system_prompt)
582 .saturating_add(tools_bytes.div_ceil(self.config.bytes_per_token))
583 .saturating_add(self.config.reserve_output_tokens)
584 .saturating_add(self.config.safety_margin_tokens);
585 if static_tokens >= self.config.max_tokens {
586 return Err(ContextError(
587 "system prompt and tool schemas exceed context budget".into(),
588 ));
589 }
590 let budget = self.config.max_tokens - static_tokens;
591 let groups = Self::group_messages(history);
592 let newest = groups.last().expect("history produced at least one group");
593 let newest_cost: usize = newest
594 .iter()
595 .map(|message| message.estimated_tokens(self.config.bytes_per_token))
596 .sum();
597 if newest_cost > budget {
598 return Err(ContextError("newest turn exceeds context budget".into()));
599 }
600
601 let before_history_tokens: usize = history
602 .iter()
603 .map(|message| message.estimated_tokens(self.config.bytes_per_token))
604 .sum();
605 let mut selected_groups: Vec<Vec<Message>> = vec![newest.clone()];
606 let mut selected_cost = newest_cost;
607 for group in groups[..groups.len() - 1].iter().rev() {
608 let cost: usize = group
609 .iter()
610 .map(|message| message.estimated_tokens(self.config.bytes_per_token))
611 .sum();
612 if selected_cost.saturating_add(cost) <= budget {
613 selected_groups.insert(0, group.clone());
614 selected_cost += cost;
615 } else {
616 break;
617 }
618 }
619
620 let mut removed_messages = groups[..groups.len() - selected_groups.len()]
621 .iter()
622 .map(Vec::len)
623 .sum::<usize>();
624 if removed_messages > 0 {
625 loop {
626 let note = Message::HistoryNote {
627 content: self.summarize(&history[..removed_messages]),
628 };
629 let note_cost = note.estimated_tokens(self.config.bytes_per_token);
630 if selected_cost.saturating_add(note_cost) <= budget {
631 let mut selected: Vec<Message> =
632 selected_groups.into_iter().flatten().collect();
633 selected.insert(0, note);
634 selected_cost += note_cost;
635 return Ok(ContextSelection {
636 messages: selected,
637 before_tokens: static_tokens.saturating_add(before_history_tokens),
638 after_tokens: static_tokens.saturating_add(selected_cost),
639 removed_messages,
640 });
641 }
642 if selected_groups.len() == 1 {
643 let available_tokens = budget.saturating_sub(selected_cost);
644 let content = match note {
645 Message::HistoryNote { content } => content,
646 _ => unreachable!(),
647 };
648 let Some(note) =
649 fit_history_note(&content, available_tokens, self.config.bytes_per_token)
650 else {
651 return Err(ContextError(
652 "compaction note cannot fit context budget".into(),
653 ));
654 };
655 let note_cost = note.estimated_tokens(self.config.bytes_per_token);
656 let mut selected: Vec<Message> =
657 selected_groups.into_iter().flatten().collect();
658 selected.insert(0, note);
659 selected_cost += note_cost;
660 return Ok(ContextSelection {
661 messages: selected,
662 before_tokens: static_tokens.saturating_add(before_history_tokens),
663 after_tokens: static_tokens.saturating_add(selected_cost),
664 removed_messages,
665 });
666 }
667 let removed_group = selected_groups.remove(0);
668 let removed_cost: usize = removed_group
669 .iter()
670 .map(|message| message.estimated_tokens(self.config.bytes_per_token))
671 .sum();
672 selected_cost = selected_cost.saturating_sub(removed_cost);
673 removed_messages += removed_group.len();
674 }
675 }
676
677 let selected: Vec<Message> = selected_groups.into_iter().flatten().collect();
678 Ok(ContextSelection {
679 messages: selected,
680 before_tokens: static_tokens.saturating_add(before_history_tokens),
681 after_tokens: static_tokens.saturating_add(selected_cost),
682 removed_messages,
683 })
684 }
685}
686
687#[derive(Debug, Clone)]
688pub struct HistoryLimits {
689 pub max_bytes: usize,
690 pub max_messages: usize,
691 pub note_max_chars: usize,
692}
693
694impl Default for HistoryLimits {
695 fn default() -> Self {
696 Self {
697 max_bytes: 16 * 1024 * 1024,
698 max_messages: 10_000,
699 note_max_chars: 4_000,
700 }
701 }
702}
703
704#[derive(Debug, Clone)]
705pub struct AgentConfig {
706 pub system_prompt: String,
707 pub max_steps: usize,
708 pub history_limits: HistoryLimits,
709}
710
711#[derive(Debug, Clone)]
712pub struct TurnOutcome {
713 pub steps: usize,
714 pub usage: Usage,
715}
716
717#[derive(Debug, Error)]
718pub enum AgentError {
719 #[error("turn cancelled")]
720 Cancelled,
721 #[error("{0}")]
722 Provider(String),
723 #[error("{0}")]
724 ContextLimit(String),
725 #[error("agent reached its maximum step count")]
726 StepLimit,
727 #[error("{0}")]
728 HistoryLimit(String),
729 #[error("{0}")]
730 ResponseLimit(String),
731 #[error("{0}")]
732 ToolLimit(String),
733 #[error("{0}")]
734 Internal(String),
735}
736
737impl AgentError {
738 pub fn code(&self) -> &'static str {
739 match self {
740 Self::Cancelled => "cancelled",
741 Self::Provider(_) => "provider_error",
742 Self::ContextLimit(_) => "context_limit",
743 Self::StepLimit => "step_limit",
744 Self::HistoryLimit(_) => "history_limit",
745 Self::ResponseLimit(_) => "response_limit",
746 Self::ToolLimit(_) => "tool_limit",
747 Self::Internal(_) => "internal_error",
748 }
749 }
750}
751
752pub struct AgentRuntime {
753 provider: Arc<dyn Provider>,
754 tools: Arc<ToolRegistry>,
755 context: Arc<dyn ContextPolicy>,
756 config: AgentConfig,
757 workspace: PathBuf,
758}
759
760impl AgentRuntime {
761 pub fn new(
762 provider: Arc<dyn Provider>,
763 tools: Arc<ToolRegistry>,
764 context: Arc<dyn ContextPolicy>,
765 config: AgentConfig,
766 workspace: PathBuf,
767 ) -> Self {
768 Self {
769 provider,
770 tools,
771 context,
772 config,
773 workspace,
774 }
775 }
776
777 pub fn model(&self) -> &str {
778 self.provider.model()
779 }
780
781 pub async fn run_turn(
782 &self,
783 history: &mut Vec<Message>,
784 prompt: String,
785 sink: Arc<dyn EventSink>,
786 approvals: Arc<dyn ApprovalGate>,
787 cancellation: CancellationToken,
788 ) -> Result<TurnOutcome, AgentError> {
789 let checkpoint = history.clone();
790 let result = self
791 .run_turn_inner(history, prompt, sink, approvals, cancellation)
792 .await;
793 if result.is_err() {
794 *history = checkpoint;
795 }
796 result
797 }
798
799 async fn run_turn_inner(
800 &self,
801 history: &mut Vec<Message>,
802 prompt: String,
803 sink: Arc<dyn EventSink>,
804 approvals: Arc<dyn ApprovalGate>,
805 cancellation: CancellationToken,
806 ) -> Result<TurnOutcome, AgentError> {
807 if cancellation.is_cancelled() {
808 return Err(AgentError::Cancelled);
809 }
810 history.push(Message::User { content: prompt });
811 self.enforce_history_limits(history, sink.as_ref()).await?;
812 let specs = self.tools.specs();
813 let mut usage = Usage::default();
814
815 for step in 1..=self.config.max_steps {
816 if cancellation.is_cancelled() {
817 return Err(AgentError::Cancelled);
818 }
819 let selection = self
820 .context
821 .select(history, &self.config.system_prompt, &specs)
822 .map_err(|error| AgentError::ContextLimit(error.to_string()))?;
823 if selection.removed_messages > 0 {
824 sink.emit(CoreEvent::ContextCompacted {
825 before_tokens: selection.before_tokens,
826 after_tokens: selection.after_tokens,
827 removed_messages: selection.removed_messages,
828 })
829 .await?;
830 }
831 let delta_sink: Arc<dyn TextDeltaSink> = Arc::new(ForwardDeltas {
832 sink: Arc::clone(&sink),
833 });
834 let response = self
835 .provider
836 .complete(
837 ProviderRequest {
838 system_prompt: self.config.system_prompt.clone(),
839 messages: selection.messages,
840 tools: specs.clone(),
841 },
842 delta_sink,
843 cancellation.child_token(),
844 )
845 .await
846 .map_err(map_provider_error)?;
847 usage.add(&response.usage);
848 sink.emit(CoreEvent::AssistantCompleted {
849 content: response.content.clone(),
850 })
851 .await?;
852 let calls = response.tool_calls.clone();
853 history.push(Message::Assistant {
854 content: response.content,
855 tool_calls: response.tool_calls,
856 });
857 self.enforce_history_limits(history, sink.as_ref()).await?;
858 if calls.is_empty() {
859 return Ok(TurnOutcome { steps: step, usage });
860 }
861
862 for call in calls {
863 if cancellation.is_cancelled() {
864 return Err(AgentError::Cancelled);
865 }
866 sink.emit(CoreEvent::ToolProposed {
867 call_id: call.id.clone(),
868 name: call.name.clone(),
869 arguments: call.arguments.clone(),
870 })
871 .await?;
872 let Some(tool) = self.tools.get(&call.name) else {
873 let output = ToolOutput::failure(format!("unknown tool: {}", call.name));
874 sink.emit(CoreEvent::ToolCompleted {
875 call_id: call.id.clone(),
876 name: call.name.clone(),
877 output: output.clone(),
878 })
879 .await?;
880 history.push(Message::Tool {
881 call_id: call.id,
882 name: call.name,
883 content: output.content,
884 is_error: true,
885 });
886 self.enforce_history_limits(history, sink.as_ref()).await?;
887 continue;
888 };
889 let risk = match tool.risk(&call.arguments) {
890 Ok(risk) => risk,
891 Err(error) => {
892 self.record_tool_error(history, sink.as_ref(), &call, error.to_string())
893 .await?;
894 self.enforce_history_limits(history, sink.as_ref()).await?;
895 continue;
896 }
897 };
898 let summary = match tool.approval_summary(&call.arguments) {
899 Ok(summary) => summary,
900 Err(error) => {
901 self.record_tool_error(history, sink.as_ref(), &call, error.to_string())
902 .await?;
903 self.enforce_history_limits(history, sink.as_ref()).await?;
904 continue;
905 }
906 };
907 let approved = approvals
908 .approve(
909 ApprovalRequest {
910 call_id: call.id.clone(),
911 name: call.name.clone(),
912 risk,
913 cwd: self.workspace.clone(),
914 summary,
915 },
916 cancellation.child_token(),
917 )
918 .await?;
919 let output = if approved {
920 sink.emit(CoreEvent::ToolStarted {
921 call_id: call.id.clone(),
922 name: call.name.clone(),
923 })
924 .await?;
925 let progress = ProgressSink::buffered();
926 let execution = tool.execute(
927 call.arguments.clone(),
928 ToolContext {
929 workspace: self.workspace.clone(),
930 cancellation: cancellation.child_token(),
931 progress: progress.clone(),
932 },
933 );
934 forward_progress(execution, &progress, sink.as_ref(), &call.id)
935 .await
936 .unwrap_or_else(|error| ToolOutput::failure(error.to_string()))
937 } else {
938 ToolOutput::failure("tool call denied by policy or user")
939 };
940 sink.emit(CoreEvent::ToolCompleted {
941 call_id: call.id.clone(),
942 name: call.name.clone(),
943 output: output.clone(),
944 })
945 .await?;
946 history.push(Message::Tool {
947 call_id: call.id,
948 name: call.name,
949 content: output.content,
950 is_error: output.is_error,
951 });
952 self.enforce_history_limits(history, sink.as_ref()).await?;
953 }
954 }
955 Err(AgentError::StepLimit)
956 }
957
958 async fn record_tool_error(
959 &self,
960 history: &mut Vec<Message>,
961 sink: &dyn EventSink,
962 call: &ToolCall,
963 message: String,
964 ) -> Result<(), AgentError> {
965 let output = ToolOutput::failure(message);
966 sink.emit(CoreEvent::ToolCompleted {
967 call_id: call.id.clone(),
968 name: call.name.clone(),
969 output: output.clone(),
970 })
971 .await?;
972 history.push(Message::Tool {
973 call_id: call.id.clone(),
974 name: call.name.clone(),
975 content: output.content,
976 is_error: true,
977 });
978 Ok(())
979 }
980
981 async fn enforce_history_limits(
982 &self,
983 history: &mut Vec<Message>,
984 sink: &dyn EventSink,
985 ) -> Result<(), AgentError> {
986 let limits = &self.config.history_limits;
987 let mut total_removed = 0;
988 while history.len() > limits.max_messages || history_bytes(history) > limits.max_bytes {
989 let latest_user = history
990 .iter()
991 .rposition(|message| matches!(message, Message::User { .. }))
992 .unwrap_or(0);
993 let active = &history[latest_user..];
994 if active.len() > limits.max_messages || history_bytes(active) > limits.max_bytes {
995 return Err(AgentError::HistoryLimit(
996 "active turn exceeds configured session history limit".into(),
997 ));
998 }
999 let first_user = history
1000 .iter()
1001 .position(|message| matches!(message, Message::User { .. }))
1002 .unwrap_or(latest_user);
1003 if first_user == latest_user {
1004 if matches!(history.first(), Some(Message::HistoryNote { .. })) {
1005 history.remove(0);
1006 total_removed += 1;
1007 continue;
1008 }
1009 return Err(AgentError::HistoryLimit(
1010 "session history cannot be reduced within its configured limit".into(),
1011 ));
1012 }
1013 let end = history[first_user + 1..]
1014 .iter()
1015 .position(|message| matches!(message, Message::User { .. }))
1016 .map(|index| first_user + 1 + index)
1017 .ok_or_else(|| {
1018 AgentError::HistoryLimit(
1019 "session history has no complete group available to trim".into(),
1020 )
1021 })?;
1022 let removed: Vec<Message> = history.drain(..end).collect();
1023 total_removed += removed.len();
1024 let note = Message::HistoryNote {
1025 content: summarize_history_trim(&removed, total_removed, limits.note_max_chars),
1026 };
1027 if matches!(history.first(), Some(Message::HistoryNote { .. })) {
1028 history.remove(0);
1029 }
1030 history.insert(0, note);
1031 }
1032 if total_removed > 0 {
1033 sink.emit(CoreEvent::SessionTrimmed {
1034 removed_messages: total_removed,
1035 history_bytes: history_bytes(history),
1036 })
1037 .await?;
1038 }
1039 Ok(())
1040 }
1041}
1042
1043async fn forward_progress<T>(
1047 execution: impl Future<Output = T>,
1048 progress: &ProgressSink,
1049 sink: &dyn EventSink,
1050 call_id: &str,
1051) -> T {
1052 let mut execution = std::pin::pin!(execution);
1053 let mut ticker = tokio::time::interval_at(
1054 tokio::time::Instant::now() + PROGRESS_INTERVAL,
1055 PROGRESS_INTERVAL,
1056 );
1057 ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
1058 let mut last_sent: Option<tokio::time::Instant> = None;
1059 let mut forwarding = true;
1060 let result = loop {
1061 tokio::select! {
1062 biased;
1063 result = &mut execution => break result,
1064 _ = ticker.tick(), if forwarding => {
1065 if let Some(text) = progress.take() {
1066 let event = CoreEvent::ToolProgress { call_id: call_id.to_owned(), text };
1067 forwarding = sink.emit(event).await.is_ok();
1068 last_sent = Some(tokio::time::Instant::now());
1069 }
1070 }
1071 }
1072 };
1073 if forwarding
1075 && last_sent.is_none_or(|sent| sent.elapsed() >= PROGRESS_INTERVAL)
1076 && let Some(text) = progress.take()
1077 {
1078 let _ = sink
1079 .emit(CoreEvent::ToolProgress {
1080 call_id: call_id.to_owned(),
1081 text,
1082 })
1083 .await;
1084 }
1085 result
1086}
1087
1088fn map_provider_error(error: ProviderError) -> AgentError {
1089 match error.kind {
1090 ProviderErrorKind::Provider => AgentError::Provider(error.message),
1091 ProviderErrorKind::ResponseLimit => AgentError::ResponseLimit(error.message),
1092 ProviderErrorKind::ToolLimit => AgentError::ToolLimit(error.message),
1093 ProviderErrorKind::Cancelled => AgentError::Cancelled,
1094 }
1095}
1096
1097struct ForwardDeltas {
1098 sink: Arc<dyn EventSink>,
1099}
1100
1101#[async_trait]
1102impl TextDeltaSink for ForwardDeltas {
1103 async fn push(&self, delta: &str) -> Result<(), ProviderError> {
1104 self.sink
1105 .emit(CoreEvent::AssistantDelta {
1106 content: delta.to_owned(),
1107 })
1108 .await
1109 .map_err(|error| match error {
1110 AgentError::Cancelled => {
1111 ProviderError::new(ProviderErrorKind::Cancelled, "turn cancelled")
1112 }
1113 AgentError::ResponseLimit(message) => {
1114 ProviderError::new(ProviderErrorKind::ResponseLimit, message)
1115 }
1116 AgentError::ToolLimit(message) => {
1117 ProviderError::new(ProviderErrorKind::ToolLimit, message)
1118 }
1119 error => ProviderError::new(ProviderErrorKind::Provider, error.to_string()),
1120 })
1121 }
1122}
1123
1124fn history_bytes(history: &[Message]) -> usize {
1125 serde_json::to_vec(history).map_or(usize::MAX, |value| value.len())
1126}
1127
1128fn summarize_history_trim(messages: &[Message], removed: usize, max_chars: usize) -> String {
1129 let mut note =
1130 format!("[SCV trimmed {removed} earlier canonical messages to enforce session limits.]\n");
1131 for message in messages {
1132 let (label, content) = match message {
1133 Message::User { content } => ("user", content.as_str()),
1134 Message::Assistant { content, .. } => ("assistant", content.as_str()),
1135 Message::Tool {
1136 name,
1137 content,
1138 is_error,
1139 ..
1140 } => {
1141 let status = if *is_error { "failed" } else { "ok" };
1142 note.push_str(&format!("tool {name} ({status}): "));
1143 ("", content.as_str())
1144 }
1145 Message::HistoryNote { content } => ("earlier", content.as_str()),
1146 };
1147 if !label.is_empty() {
1148 note.push_str(label);
1149 note.push_str(": ");
1150 }
1151 note.push_str(&char_tail(content, 160).replace('\n', " "));
1152 note.push('\n');
1153 if note.chars().count() >= max_chars {
1154 break;
1155 }
1156 }
1157 truncate_chars(¬e, max_chars)
1158}
1159
1160fn truncate_chars(value: &str, max_chars: usize) -> String {
1161 value.chars().take(max_chars).collect()
1162}
1163
1164fn fit_history_note(
1165 content: &str,
1166 available_tokens: usize,
1167 bytes_per_token: usize,
1168) -> Option<Message> {
1169 let chars: Vec<char> = content.chars().collect();
1170 let mut low = 0usize;
1171 let mut high = chars.len();
1172 let mut best = None;
1173 while low <= high {
1174 let middle = low + (high - low) / 2;
1175 let candidate = Message::HistoryNote {
1176 content: chars[..middle].iter().collect(),
1177 };
1178 if candidate.estimated_tokens(bytes_per_token) <= available_tokens {
1179 best = Some(candidate);
1180 low = middle.saturating_add(1);
1181 } else if middle == 0 {
1182 break;
1183 } else {
1184 high = middle - 1;
1185 }
1186 }
1187 best
1188}
1189
1190fn char_tail(value: &str, max_chars: usize) -> String {
1191 let count = value.chars().count();
1192 value
1193 .chars()
1194 .skip(count.saturating_sub(max_chars))
1195 .collect()
1196}
1197
1198#[cfg(test)]
1199mod tests {
1200 use std::{collections::VecDeque, sync::Mutex};
1201
1202 use tokio::sync::Notify;
1203
1204 use super::*;
1205
1206 struct ScriptedProvider {
1207 responses: Mutex<VecDeque<AssistantResponse>>,
1208 }
1209
1210 #[async_trait]
1211 impl Provider for ScriptedProvider {
1212 fn model(&self) -> &str {
1213 "test-model"
1214 }
1215
1216 async fn complete(
1217 &self,
1218 _request: ProviderRequest,
1219 deltas: Arc<dyn TextDeltaSink>,
1220 _cancellation: CancellationToken,
1221 ) -> Result<AssistantResponse, ProviderError> {
1222 let response = self.responses.lock().unwrap().pop_front().unwrap();
1223 deltas.push(&response.content).await?;
1224 Ok(response)
1225 }
1226 }
1227
1228 struct CollectSink(Mutex<Vec<CoreEvent>>);
1229
1230 #[async_trait]
1231 impl EventSink for CollectSink {
1232 async fn emit(&self, event: CoreEvent) -> Result<(), AgentError> {
1233 self.0.lock().unwrap().push(event);
1234 Ok(())
1235 }
1236 }
1237
1238 struct Allow;
1239
1240 #[async_trait]
1241 impl ApprovalGate for Allow {
1242 async fn approve(
1243 &self,
1244 _request: ApprovalRequest,
1245 _cancellation: CancellationToken,
1246 ) -> Result<bool, AgentError> {
1247 Ok(true)
1248 }
1249 }
1250
1251 struct Deny;
1252
1253 #[async_trait]
1254 impl ApprovalGate for Deny {
1255 async fn approve(
1256 &self,
1257 _request: ApprovalRequest,
1258 _cancellation: CancellationToken,
1259 ) -> Result<bool, AgentError> {
1260 Ok(false)
1261 }
1262 }
1263
1264 struct WaitForCancellation {
1265 entered: Arc<Notify>,
1266 }
1267
1268 #[async_trait]
1269 impl ApprovalGate for WaitForCancellation {
1270 async fn approve(
1271 &self,
1272 _request: ApprovalRequest,
1273 cancellation: CancellationToken,
1274 ) -> Result<bool, AgentError> {
1275 self.entered.notify_one();
1276 cancellation.cancelled().await;
1277 Err(AgentError::Cancelled)
1278 }
1279 }
1280
1281 struct EchoTool;
1282
1283 #[async_trait]
1284 impl Tool for EchoTool {
1285 fn spec(&self) -> ToolSpec {
1286 ToolSpec {
1287 name: "echo".into(),
1288 description: "Echo a value".into(),
1289 parameters: serde_json::json!({
1290 "type":"object",
1291 "properties":{"value":{"type":"string"}},
1292 "required":["value"]
1293 }),
1294 }
1295 }
1296
1297 fn risk(&self, arguments: &Value) -> Result<ToolRisk, ToolError> {
1298 arguments
1299 .get("value")
1300 .and_then(Value::as_str)
1301 .ok_or_else(|| ToolError("value must be a string".into()))?;
1302 Ok(ToolRisk::ReadOnly)
1303 }
1304
1305 fn approval_summary(&self, arguments: &Value) -> Result<String, ToolError> {
1306 self.risk(arguments)?;
1307 Ok("Echo a value".into())
1308 }
1309
1310 async fn execute(
1311 &self,
1312 arguments: Value,
1313 _context: ToolContext,
1314 ) -> Result<ToolOutput, ToolError> {
1315 Ok(ToolOutput::success(
1316 arguments["value"].as_str().unwrap_or_default(),
1317 ))
1318 }
1319 }
1320
1321 #[test]
1322 fn progress_lines_are_single_bounded_lines() {
1323 assert_eq!(
1324 progress_line(" run\n\tcargo \u{7}test "),
1325 "run cargo test"
1326 );
1327 let long = progress_line(&"x".repeat(1000));
1328 assert!(long.len() <= MAX_PROGRESS_LINE_BYTES && long.ends_with(PROGRESS_ELIDED));
1329 let wide = progress_line(&"é".repeat(300));
1330 assert!(wide.len() <= MAX_PROGRESS_LINE_BYTES);
1331 let discard = ProgressSink::default();
1332 discard.report("ignored");
1333 assert!(!discard.is_enabled() && discard.take().is_none());
1334 }
1335
1336 #[test]
1337 fn progress_events_keep_the_newest_lines_within_the_limit() {
1338 let progress = ProgressSink::buffered();
1339 assert!(progress.take().is_none());
1340 progress.report("first");
1341 progress.report("second");
1342 assert_eq!(progress.take().as_deref(), Some("first\nsecond"));
1343 assert!(progress.take().is_none());
1344 for index in 0..50 {
1345 progress.report(&format!("{index:03} {}", "y".repeat(96)));
1346 }
1347 let text = progress.take().unwrap();
1348 assert!(text.len() <= MAX_PROGRESS_EVENT_BYTES, "{}", text.len());
1349 assert!(text.starts_with(&format!("{PROGRESS_ELIDED}\n")));
1350 assert!(text.lines().last().unwrap().starts_with("049 "));
1351 progress.report("after");
1352 assert_eq!(progress.take().as_deref(), Some("after"));
1353 }
1354
1355 struct ProgressTool;
1356
1357 #[async_trait]
1358 impl Tool for ProgressTool {
1359 fn spec(&self) -> ToolSpec {
1360 ToolSpec {
1361 name: "work".into(),
1362 description: "Report progress while working".into(),
1363 parameters: serde_json::json!({"type":"object"}),
1364 }
1365 }
1366
1367 fn risk(&self, _arguments: &Value) -> Result<ToolRisk, ToolError> {
1368 Ok(ToolRisk::ReadOnly)
1369 }
1370
1371 fn approval_summary(&self, _arguments: &Value) -> Result<String, ToolError> {
1372 Ok("Work".into())
1373 }
1374
1375 async fn execute(
1376 &self,
1377 _arguments: Value,
1378 context: ToolContext,
1379 ) -> Result<ToolOutput, ToolError> {
1380 for step in 0..22 {
1381 context.progress.report(&format!("step {step}"));
1382 tokio::time::sleep(Duration::from_millis(100)).await;
1383 }
1384 Ok(ToolOutput::success("worked"))
1385 }
1386 }
1387
1388 struct TimedSink(Mutex<Vec<(tokio::time::Instant, CoreEvent)>>);
1389
1390 #[async_trait]
1391 impl EventSink for TimedSink {
1392 async fn emit(&self, event: CoreEvent) -> Result<(), AgentError> {
1393 self.0
1394 .lock()
1395 .unwrap()
1396 .push((tokio::time::Instant::now(), event));
1397 Ok(())
1398 }
1399 }
1400
1401 #[tokio::test(start_paused = true)]
1402 async fn tool_progress_is_paced_and_kept_out_of_history() {
1403 let provider = Arc::new(ScriptedProvider {
1404 responses: Mutex::new(VecDeque::from([
1405 AssistantResponse {
1406 content: String::new(),
1407 tool_calls: vec![ToolCall {
1408 id: "call-1".into(),
1409 name: "work".into(),
1410 arguments: serde_json::json!({}),
1411 }],
1412 usage: Usage::default(),
1413 },
1414 AssistantResponse {
1415 content: "done".into(),
1416 tool_calls: Vec::new(),
1417 usage: Usage::default(),
1418 },
1419 ])),
1420 });
1421 let mut registry = ToolRegistry::default();
1422 registry.register(Arc::new(ProgressTool)).unwrap();
1423 let runtime = AgentRuntime::new(
1424 provider,
1425 Arc::new(registry),
1426 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1427 AgentConfig {
1428 system_prompt: "test".into(),
1429 max_steps: 3,
1430 history_limits: HistoryLimits::default(),
1431 },
1432 PathBuf::from("/tmp"),
1433 );
1434 let sink = Arc::new(TimedSink(Mutex::new(Vec::new())));
1435 let mut history = Vec::new();
1436 runtime
1437 .run_turn(
1438 &mut history,
1439 "go".into(),
1440 sink.clone(),
1441 Arc::new(Allow),
1442 CancellationToken::new(),
1443 )
1444 .await
1445 .unwrap();
1446 let events = sink.0.lock().unwrap();
1447 let started = events
1448 .iter()
1449 .position(|(_, event)| matches!(event, CoreEvent::ToolStarted { .. }))
1450 .unwrap();
1451 let completed = events
1452 .iter()
1453 .position(|(_, event)| matches!(event, CoreEvent::ToolCompleted { .. }))
1454 .unwrap();
1455 let progress: Vec<_> = events
1456 .iter()
1457 .enumerate()
1458 .filter_map(|(index, (at, event))| match event {
1459 CoreEvent::ToolProgress { call_id, text } => Some((index, *at, call_id, text)),
1460 _ => None,
1461 })
1462 .collect();
1463 assert!((4..=5).contains(&progress.len()), "{}", progress.len());
1466 for (index, _, call_id, text) in &progress {
1467 assert!(*index > started && *index < completed);
1468 assert_eq!(call_id.as_str(), "call-1");
1469 assert!(text.len() <= MAX_PROGRESS_EVENT_BYTES);
1470 }
1471 for pair in progress.windows(2) {
1472 assert!(pair[1].1 - pair[0].1 >= PROGRESS_INTERVAL);
1473 }
1474 let all: Vec<&str> = progress
1475 .iter()
1476 .flat_map(|(_, _, _, text)| text.lines())
1477 .collect();
1478 assert_eq!(all.first(), Some(&"step 0"));
1479 assert!(all.contains(&"step 19"));
1482 let stored = serde_json::to_string(&history).unwrap();
1483 assert!(!stored.contains("step 1"), "progress leaked into history");
1484 }
1485
1486 #[tokio::test]
1487 async fn completes_a_simple_turn() {
1488 let provider = Arc::new(ScriptedProvider {
1489 responses: Mutex::new(VecDeque::from([AssistantResponse {
1490 content: "done".into(),
1491 tool_calls: Vec::new(),
1492 usage: Usage {
1493 input_tokens: Some(3),
1494 output_tokens: Some(1),
1495 },
1496 }])),
1497 });
1498 let runtime = AgentRuntime::new(
1499 provider,
1500 Arc::new(ToolRegistry::default()),
1501 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1502 AgentConfig {
1503 system_prompt: "test".into(),
1504 max_steps: 2,
1505 history_limits: HistoryLimits::default(),
1506 },
1507 PathBuf::from("/tmp"),
1508 );
1509 let sink = Arc::new(CollectSink(Mutex::new(Vec::new())));
1510 let mut history = Vec::new();
1511 let outcome = runtime
1512 .run_turn(
1513 &mut history,
1514 "hello".into(),
1515 sink.clone(),
1516 Arc::new(Allow),
1517 CancellationToken::new(),
1518 )
1519 .await
1520 .unwrap();
1521 assert_eq!(outcome.steps, 1);
1522 assert_eq!(history.len(), 2);
1523 assert!(matches!(
1524 sink.0.lock().unwrap().last(),
1525 Some(CoreEvent::AssistantCompleted { .. })
1526 ));
1527 }
1528
1529 #[test]
1530 fn context_keeps_tool_groups_together() {
1531 let policy = BudgetContextPolicy::new(ContextConfig {
1532 max_tokens: 120,
1533 reserve_output_tokens: 10,
1534 safety_margin_tokens: 10,
1535 bytes_per_token: 3,
1536 summary_max_chars: 120,
1537 })
1538 .unwrap();
1539 let history = vec![
1540 Message::User {
1541 content: "old request ".repeat(20),
1542 },
1543 Message::Assistant {
1544 content: String::new(),
1545 tool_calls: vec![ToolCall {
1546 id: "1".into(),
1547 name: "read".into(),
1548 arguments: serde_json::json!({"path":"a"}),
1549 }],
1550 },
1551 Message::Tool {
1552 call_id: "1".into(),
1553 name: "read".into(),
1554 content: "result".into(),
1555 is_error: false,
1556 },
1557 Message::User {
1558 content: "new".into(),
1559 },
1560 ];
1561 let selection = policy.select(&history, "system", &[]).unwrap();
1562 assert!(selection.removed_messages > 0);
1563 assert_eq!(
1564 selection.removed_messages,
1565 history.len() - (selection.messages.len() - 1)
1566 );
1567 assert!(matches!(
1568 selection.messages.last(),
1569 Some(Message::User { .. })
1570 ));
1571 assert!(
1572 !selection
1573 .messages
1574 .iter()
1575 .any(|message| matches!(message, Message::Tool { call_id, .. } if call_id == "1"))
1576 );
1577 }
1578
1579 #[tokio::test]
1580 async fn repeated_history_trimming_rebuilds_the_note_and_makes_progress() {
1581 let runtime = AgentRuntime::new(
1582 Arc::new(ScriptedProvider {
1583 responses: Mutex::new(VecDeque::new()),
1584 }),
1585 Arc::new(ToolRegistry::default()),
1586 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1587 AgentConfig {
1588 system_prompt: "test".into(),
1589 max_steps: 1,
1590 history_limits: HistoryLimits {
1591 max_bytes: 4096,
1592 max_messages: 3,
1593 note_max_chars: 80,
1594 },
1595 },
1596 PathBuf::from("/tmp"),
1597 );
1598 let sink = CollectSink(Mutex::new(Vec::new()));
1599 let mut history = vec![
1600 Message::HistoryNote {
1601 content: "previous trim".into(),
1602 },
1603 Message::User {
1604 content: "old request".into(),
1605 },
1606 Message::Assistant {
1607 content: "old answer".into(),
1608 tool_calls: Vec::new(),
1609 },
1610 Message::User {
1611 content: "active request".into(),
1612 },
1613 ];
1614 runtime
1615 .enforce_history_limits(&mut history, &sink)
1616 .await
1617 .unwrap();
1618 assert!(history.len() <= 3);
1619 assert!(matches!(history.first(), Some(Message::HistoryNote { .. })));
1620 assert!(matches!(history.last(), Some(Message::User { .. })));
1621 }
1622
1623 #[tokio::test]
1624 async fn active_turn_over_history_limit_rolls_back() {
1625 let runtime = AgentRuntime::new(
1626 Arc::new(ScriptedProvider {
1627 responses: Mutex::new(VecDeque::new()),
1628 }),
1629 Arc::new(ToolRegistry::default()),
1630 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1631 AgentConfig {
1632 system_prompt: "test".into(),
1633 max_steps: 1,
1634 history_limits: HistoryLimits {
1635 max_bytes: 16,
1636 max_messages: 10,
1637 note_max_chars: 8,
1638 },
1639 },
1640 PathBuf::from("/tmp"),
1641 );
1642 let sink = CollectSink(Mutex::new(Vec::new()));
1643 let mut history = Vec::new();
1644 let result = runtime
1645 .run_turn(
1646 &mut history,
1647 "too large for the configured history".into(),
1648 Arc::new(sink),
1649 Arc::new(Allow),
1650 CancellationToken::new(),
1651 )
1652 .await;
1653 assert!(matches!(result, Err(AgentError::HistoryLimit(_))));
1654 assert!(history.is_empty());
1655 }
1656
1657 #[tokio::test]
1658 async fn executes_a_multi_step_tool_loop_and_aggregates_usage() {
1659 let provider = Arc::new(ScriptedProvider {
1660 responses: Mutex::new(VecDeque::from([
1661 AssistantResponse {
1662 content: String::new(),
1663 tool_calls: vec![ToolCall {
1664 id: "call-1".into(),
1665 name: "echo".into(),
1666 arguments: serde_json::json!({"value":"hello"}),
1667 }],
1668 usage: Usage {
1669 input_tokens: Some(2),
1670 output_tokens: Some(1),
1671 },
1672 },
1673 AssistantResponse {
1674 content: "done".into(),
1675 tool_calls: Vec::new(),
1676 usage: Usage {
1677 input_tokens: Some(4),
1678 output_tokens: Some(2),
1679 },
1680 },
1681 ])),
1682 });
1683 let mut registry = ToolRegistry::default();
1684 registry.register(Arc::new(EchoTool)).unwrap();
1685 let runtime = AgentRuntime::new(
1686 provider,
1687 Arc::new(registry),
1688 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1689 AgentConfig {
1690 system_prompt: "test".into(),
1691 max_steps: 3,
1692 history_limits: HistoryLimits::default(),
1693 },
1694 PathBuf::from("/tmp"),
1695 );
1696 let sink = Arc::new(CollectSink(Mutex::new(Vec::new())));
1697 let mut history = Vec::new();
1698 let outcome = runtime
1699 .run_turn(
1700 &mut history,
1701 "start".into(),
1702 sink,
1703 Arc::new(Allow),
1704 CancellationToken::new(),
1705 )
1706 .await
1707 .unwrap();
1708 assert_eq!(outcome.steps, 2);
1709 assert_eq!(outcome.usage.input_tokens, Some(6));
1710 assert_eq!(outcome.usage.output_tokens, Some(3));
1711 assert!(matches!(
1712 history.get(2),
1713 Some(Message::Tool {
1714 content,
1715 is_error: false,
1716 ..
1717 }) if content == "hello"
1718 ));
1719 }
1720
1721 #[tokio::test]
1722 async fn denial_is_recorded_as_a_model_visible_tool_failure() {
1723 let provider = Arc::new(ScriptedProvider {
1724 responses: Mutex::new(VecDeque::from([
1725 AssistantResponse {
1726 content: String::new(),
1727 tool_calls: vec![ToolCall {
1728 id: "call-1".into(),
1729 name: "echo".into(),
1730 arguments: serde_json::json!({"value":"blocked"}),
1731 }],
1732 usage: Usage::default(),
1733 },
1734 AssistantResponse {
1735 content: "handled".into(),
1736 tool_calls: Vec::new(),
1737 usage: Usage::default(),
1738 },
1739 ])),
1740 });
1741 let mut registry = ToolRegistry::default();
1742 registry.register(Arc::new(EchoTool)).unwrap();
1743 let runtime = AgentRuntime::new(
1744 provider,
1745 Arc::new(registry),
1746 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1747 AgentConfig {
1748 system_prompt: "test".into(),
1749 max_steps: 3,
1750 history_limits: HistoryLimits::default(),
1751 },
1752 PathBuf::from("/tmp"),
1753 );
1754 let mut history = Vec::new();
1755 runtime
1756 .run_turn(
1757 &mut history,
1758 "start".into(),
1759 Arc::new(CollectSink(Mutex::new(Vec::new()))),
1760 Arc::new(Deny),
1761 CancellationToken::new(),
1762 )
1763 .await
1764 .unwrap();
1765 assert!(matches!(
1766 history.get(2),
1767 Some(Message::Tool {
1768 content,
1769 is_error: true,
1770 ..
1771 }) if content.contains("denied")
1772 ));
1773 }
1774
1775 #[tokio::test]
1776 async fn cancellation_during_approval_rolls_back_the_active_tool_group() {
1777 let provider = Arc::new(ScriptedProvider {
1778 responses: Mutex::new(VecDeque::from([AssistantResponse {
1779 content: String::new(),
1780 tool_calls: vec![ToolCall {
1781 id: "call-cancel".into(),
1782 name: "echo".into(),
1783 arguments: serde_json::json!({"value":"hello"}),
1784 }],
1785 usage: Usage::default(),
1786 }])),
1787 });
1788 let mut registry = ToolRegistry::default();
1789 registry.register(Arc::new(EchoTool)).unwrap();
1790 let runtime = AgentRuntime::new(
1791 provider,
1792 Arc::new(registry),
1793 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1794 AgentConfig {
1795 system_prompt: "test".into(),
1796 max_steps: 2,
1797 history_limits: HistoryLimits::default(),
1798 },
1799 PathBuf::from("/tmp"),
1800 );
1801 let before = vec![
1802 Message::User {
1803 content: "previous".into(),
1804 },
1805 Message::Assistant {
1806 content: "answer".into(),
1807 tool_calls: Vec::new(),
1808 },
1809 ];
1810 let mut history = before.clone();
1811 let cancellation = CancellationToken::new();
1812 let cancel = cancellation.clone();
1813 let entered = Arc::new(Notify::new());
1814 let wait = Arc::clone(&entered);
1815 let run = runtime.run_turn(
1816 &mut history,
1817 "new turn".into(),
1818 Arc::new(CollectSink(Mutex::new(Vec::new()))),
1819 Arc::new(WaitForCancellation { entered }),
1820 cancellation,
1821 );
1822 let cancel_when_waiting = async move {
1823 wait.notified().await;
1824 cancel.cancel();
1825 };
1826 let (result, ()) = tokio::join!(run, cancel_when_waiting);
1827 assert!(matches!(result, Err(AgentError::Cancelled)));
1828 assert_eq!(history, before);
1829 }
1830
1831 #[tokio::test]
1832 async fn stops_after_the_configured_maximum_step() {
1833 let provider = Arc::new(ScriptedProvider {
1834 responses: Mutex::new(VecDeque::from([AssistantResponse {
1835 content: String::new(),
1836 tool_calls: vec![ToolCall {
1837 id: "call-1".into(),
1838 name: "echo".into(),
1839 arguments: serde_json::json!({"value":"one"}),
1840 }],
1841 usage: Usage::default(),
1842 }])),
1843 });
1844 let mut registry = ToolRegistry::default();
1845 registry.register(Arc::new(EchoTool)).unwrap();
1846 let runtime = AgentRuntime::new(
1847 provider,
1848 Arc::new(registry),
1849 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1850 AgentConfig {
1851 system_prompt: "test".into(),
1852 max_steps: 1,
1853 history_limits: HistoryLimits::default(),
1854 },
1855 PathBuf::from("/tmp"),
1856 );
1857 let result = runtime
1858 .run_turn(
1859 &mut Vec::new(),
1860 "start".into(),
1861 Arc::new(CollectSink(Mutex::new(Vec::new()))),
1862 Arc::new(Allow),
1863 CancellationToken::new(),
1864 )
1865 .await;
1866 assert!(matches!(result, Err(AgentError::StepLimit)));
1867 }
1868
1869 #[test]
1870 fn duplicate_tool_registration_does_not_replace_the_original() {
1871 let mut registry = ToolRegistry::default();
1872 registry.register(Arc::new(EchoTool)).unwrap();
1873 assert!(registry.register(Arc::new(EchoTool)).is_err());
1874 assert_eq!(registry.tools.len(), 1);
1875 assert!(registry.get("echo").is_some());
1876 }
1877}