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