1use std::{collections::HashMap, path::PathBuf, sync::Arc};
4
5use async_trait::async_trait;
6use serde::{Deserialize, Serialize};
7use serde_json::Value;
8use thiserror::Error;
9use tokio_util::sync::CancellationToken;
10
11#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
12#[serde(tag = "role", rename_all = "snake_case")]
13pub enum Message {
14 User {
15 content: String,
16 },
17 Assistant {
18 content: String,
19 #[serde(default, skip_serializing_if = "Vec::is_empty")]
20 tool_calls: Vec<ToolCall>,
21 },
22 Tool {
23 call_id: String,
24 name: String,
25 content: String,
26 is_error: bool,
27 },
28 HistoryNote {
29 content: String,
30 },
31}
32
33impl Message {
34 fn estimated_tokens(&self, bytes_per_token: usize) -> usize {
35 let bytes = serde_json::to_vec(self).map_or(0, |value| value.len());
36 bytes.div_ceil(bytes_per_token).saturating_add(4)
37 }
38}
39
40#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
41pub struct ToolCall {
42 pub id: String,
43 pub name: String,
44 pub arguments: Value,
45}
46
47#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
48pub struct ToolSpec {
49 pub name: String,
50 pub description: String,
51 pub parameters: Value,
52}
53
54#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
55#[serde(rename_all = "snake_case")]
56pub enum ToolRisk {
57 ReadOnly,
58 Filesystem,
59 Process,
60 Delegate,
61 Network,
64}
65
66impl ToolRisk {
67 pub fn as_str(self) -> &'static str {
68 match self {
69 Self::ReadOnly => "read_only",
70 Self::Filesystem => "filesystem",
71 Self::Process => "process",
72 Self::Delegate => "delegate",
73 Self::Network => "network",
74 }
75 }
76}
77
78#[derive(Debug, Clone)]
79pub struct ToolContext {
80 pub workspace: PathBuf,
81 pub cancellation: CancellationToken,
82}
83
84#[derive(Debug, Clone, PartialEq, Eq)]
85pub struct ToolOutput {
86 pub content: String,
87 pub is_error: bool,
88 pub truncated: bool,
89}
90
91impl ToolOutput {
92 pub fn success(content: impl Into<String>) -> Self {
93 Self {
94 content: content.into(),
95 is_error: false,
96 truncated: false,
97 }
98 }
99
100 pub fn failure(content: impl Into<String>) -> Self {
101 Self {
102 content: content.into(),
103 is_error: true,
104 truncated: false,
105 }
106 }
107}
108
109#[derive(Debug, Error)]
110#[error("{0}")]
111pub struct ToolError(pub String);
112
113#[async_trait]
114pub trait Tool: Send + Sync {
115 fn spec(&self) -> ToolSpec;
116 fn risk(&self, arguments: &Value) -> Result<ToolRisk, ToolError>;
117 fn approval_summary(&self, arguments: &Value) -> Result<String, ToolError>;
118 async fn execute(
119 &self,
120 arguments: Value,
121 context: ToolContext,
122 ) -> Result<ToolOutput, ToolError>;
123}
124
125#[derive(Default)]
126pub struct ToolRegistry {
127 tools: HashMap<String, Arc<dyn Tool>>,
128}
129
130impl ToolRegistry {
131 pub fn register(&mut self, tool: Arc<dyn Tool>) -> Result<(), ToolError> {
132 let name = tool.spec().name;
133 if self.tools.contains_key(&name) {
134 return Err(ToolError(format!("duplicate tool name: {name}")));
135 }
136 self.tools.insert(name, tool);
137 Ok(())
138 }
139
140 pub fn get(&self, name: &str) -> Option<Arc<dyn Tool>> {
141 self.tools.get(name).cloned()
142 }
143
144 pub fn specs(&self) -> Vec<ToolSpec> {
145 let mut specs: Vec<_> = self.tools.values().map(|tool| tool.spec()).collect();
146 specs.sort_by(|a, b| a.name.cmp(&b.name));
147 specs
148 }
149}
150
151#[derive(Debug, Clone, Default, PartialEq, Eq)]
152pub struct Usage {
153 pub input_tokens: Option<u64>,
154 pub output_tokens: Option<u64>,
155}
156
157impl Usage {
158 fn add(&mut self, other: &Self) {
159 self.input_tokens = add_optional(self.input_tokens, other.input_tokens);
160 self.output_tokens = add_optional(self.output_tokens, other.output_tokens);
161 }
162}
163
164fn add_optional(left: Option<u64>, right: Option<u64>) -> Option<u64> {
165 match (left, right) {
166 (None, None) => None,
167 (left, right) => Some(left.unwrap_or(0).saturating_add(right.unwrap_or(0))),
168 }
169}
170
171#[derive(Debug, Clone)]
172pub struct ProviderRequest {
173 pub system_prompt: String,
174 pub messages: Vec<Message>,
175 pub tools: Vec<ToolSpec>,
176}
177
178#[derive(Debug, Clone)]
179pub struct AssistantResponse {
180 pub content: String,
181 pub tool_calls: Vec<ToolCall>,
182 pub usage: Usage,
183}
184
185#[derive(Debug, Clone, Copy, PartialEq, Eq)]
186pub enum ProviderErrorKind {
187 Provider,
188 ResponseLimit,
189 ToolLimit,
190 Cancelled,
191}
192
193#[derive(Debug, Error)]
194#[error("{message}")]
195pub struct ProviderError {
196 pub kind: ProviderErrorKind,
197 pub message: String,
198}
199
200impl ProviderError {
201 pub fn new(kind: ProviderErrorKind, message: impl Into<String>) -> Self {
202 Self {
203 kind,
204 message: message.into(),
205 }
206 }
207}
208
209#[async_trait]
210pub trait TextDeltaSink: Send + Sync {
211 async fn push(&self, delta: &str) -> Result<(), ProviderError>;
212}
213
214#[async_trait]
215pub trait Provider: Send + Sync {
216 fn model(&self) -> &str;
217
218 async fn complete(
219 &self,
220 request: ProviderRequest,
221 deltas: Arc<dyn TextDeltaSink>,
222 cancellation: CancellationToken,
223 ) -> Result<AssistantResponse, ProviderError>;
224}
225
226#[derive(Debug, Clone)]
227pub enum CoreEvent {
228 AssistantDelta {
229 content: String,
230 },
231 AssistantCompleted {
232 content: String,
233 },
234 ToolProposed {
235 call_id: String,
236 name: String,
237 arguments: Value,
238 },
239 ToolStarted {
240 call_id: String,
241 name: String,
242 },
243 ToolCompleted {
244 call_id: String,
245 name: String,
246 output: ToolOutput,
247 },
248 ContextCompacted {
249 before_tokens: usize,
250 after_tokens: usize,
251 removed_messages: usize,
252 },
253 SessionTrimmed {
254 removed_messages: usize,
255 history_bytes: usize,
256 },
257}
258
259#[async_trait]
260pub trait EventSink: Send + Sync {
261 async fn emit(&self, event: CoreEvent) -> Result<(), AgentError>;
262}
263
264#[derive(Debug, Clone)]
265pub struct ApprovalRequest {
266 pub call_id: String,
267 pub name: String,
268 pub risk: ToolRisk,
269 pub cwd: PathBuf,
270 pub summary: String,
271}
272
273#[async_trait]
274pub trait ApprovalGate: Send + Sync {
275 async fn approve(
276 &self,
277 request: ApprovalRequest,
278 cancellation: CancellationToken,
279 ) -> Result<bool, AgentError>;
280}
281
282#[derive(Debug, Clone)]
283pub struct ContextConfig {
284 pub max_tokens: usize,
285 pub reserve_output_tokens: usize,
286 pub safety_margin_tokens: usize,
287 pub bytes_per_token: usize,
288 pub summary_max_chars: usize,
289}
290
291impl Default for ContextConfig {
292 fn default() -> Self {
293 Self {
294 max_tokens: 128_000,
295 reserve_output_tokens: 8_192,
296 safety_margin_tokens: 2_048,
297 bytes_per_token: 3,
298 summary_max_chars: 6_000,
299 }
300 }
301}
302
303#[derive(Debug, Clone)]
304pub struct ContextSelection {
305 pub messages: Vec<Message>,
306 pub before_tokens: usize,
307 pub after_tokens: usize,
308 pub removed_messages: usize,
309}
310
311#[derive(Debug, Error)]
312#[error("{0}")]
313pub struct ContextError(pub String);
314
315pub trait ContextPolicy: Send + Sync {
316 fn select(
317 &self,
318 history: &[Message],
319 system_prompt: &str,
320 tools: &[ToolSpec],
321 ) -> Result<ContextSelection, ContextError>;
322}
323
324pub struct BudgetContextPolicy {
325 config: ContextConfig,
326}
327
328impl BudgetContextPolicy {
329 pub fn new(config: ContextConfig) -> Result<Self, ContextError> {
330 if config.bytes_per_token == 0 {
331 return Err(ContextError(
332 "context.bytes_per_token must be positive".into(),
333 ));
334 }
335 if config
336 .reserve_output_tokens
337 .saturating_add(config.safety_margin_tokens)
338 >= config.max_tokens
339 {
340 return Err(ContextError(
341 "context reserve and safety margin consume the model window".into(),
342 ));
343 }
344 Ok(Self { config })
345 }
346
347 fn string_tokens(&self, value: &str) -> usize {
348 value.len().div_ceil(self.config.bytes_per_token)
349 }
350
351 fn group_messages(history: &[Message]) -> Vec<Vec<Message>> {
352 let mut groups: Vec<Vec<Message>> = Vec::new();
353 for message in history {
354 if matches!(message, Message::User { .. }) || groups.is_empty() {
355 groups.push(Vec::new());
356 }
357 groups
358 .last_mut()
359 .expect("a group was just created")
360 .push(message.clone());
361 }
362 groups
363 }
364
365 fn summarize(&self, messages: &[Message]) -> String {
366 let mut output = format!(
367 "[SCV compacted {} earlier messages. Bounded extracts follow.]\n",
368 messages.len()
369 );
370 for message in messages {
371 let (label, content) = match message {
372 Message::User { content } => ("user", content.as_str()),
373 Message::Assistant { content, .. } => ("assistant", content.as_str()),
374 Message::Tool {
375 name,
376 content,
377 is_error,
378 ..
379 } => {
380 let status = if *is_error { "failed" } else { "ok" };
381 output.push_str(&format!("tool {name} ({status}): "));
382 ("", content.as_str())
383 }
384 Message::HistoryNote { content } => ("earlier", content.as_str()),
385 };
386 if !label.is_empty() {
387 output.push_str(label);
388 output.push_str(": ");
389 }
390 let tail = char_tail(content, 240);
391 output.push_str(&tail.replace('\n', " "));
392 output.push('\n');
393 if output.chars().count() >= self.config.summary_max_chars {
394 break;
395 }
396 }
397 truncate_chars(&output, self.config.summary_max_chars)
398 }
399}
400
401impl ContextPolicy for BudgetContextPolicy {
402 fn select(
403 &self,
404 history: &[Message],
405 system_prompt: &str,
406 tools: &[ToolSpec],
407 ) -> Result<ContextSelection, ContextError> {
408 if history.is_empty() {
409 return Ok(ContextSelection {
410 messages: Vec::new(),
411 before_tokens: 0,
412 after_tokens: 0,
413 removed_messages: 0,
414 });
415 }
416 let tools_bytes = serde_json::to_vec(tools).map_or(0, |value| value.len());
417 let static_tokens = self
418 .string_tokens(system_prompt)
419 .saturating_add(tools_bytes.div_ceil(self.config.bytes_per_token))
420 .saturating_add(self.config.reserve_output_tokens)
421 .saturating_add(self.config.safety_margin_tokens);
422 if static_tokens >= self.config.max_tokens {
423 return Err(ContextError(
424 "system prompt and tool schemas exceed context budget".into(),
425 ));
426 }
427 let budget = self.config.max_tokens - static_tokens;
428 let groups = Self::group_messages(history);
429 let newest = groups.last().expect("history produced at least one group");
430 let newest_cost: usize = newest
431 .iter()
432 .map(|message| message.estimated_tokens(self.config.bytes_per_token))
433 .sum();
434 if newest_cost > budget {
435 return Err(ContextError("newest turn exceeds context budget".into()));
436 }
437
438 let before_history_tokens: usize = history
439 .iter()
440 .map(|message| message.estimated_tokens(self.config.bytes_per_token))
441 .sum();
442 let mut selected_groups: Vec<Vec<Message>> = vec![newest.clone()];
443 let mut selected_cost = newest_cost;
444 for group in groups[..groups.len() - 1].iter().rev() {
445 let cost: usize = group
446 .iter()
447 .map(|message| message.estimated_tokens(self.config.bytes_per_token))
448 .sum();
449 if selected_cost.saturating_add(cost) <= budget {
450 selected_groups.insert(0, group.clone());
451 selected_cost += cost;
452 } else {
453 break;
454 }
455 }
456
457 let mut removed_messages = groups[..groups.len() - selected_groups.len()]
458 .iter()
459 .map(Vec::len)
460 .sum::<usize>();
461 if removed_messages > 0 {
462 loop {
463 let note = Message::HistoryNote {
464 content: self.summarize(&history[..removed_messages]),
465 };
466 let note_cost = note.estimated_tokens(self.config.bytes_per_token);
467 if selected_cost.saturating_add(note_cost) <= budget {
468 let mut selected: Vec<Message> =
469 selected_groups.into_iter().flatten().collect();
470 selected.insert(0, note);
471 selected_cost += note_cost;
472 return Ok(ContextSelection {
473 messages: selected,
474 before_tokens: static_tokens.saturating_add(before_history_tokens),
475 after_tokens: static_tokens.saturating_add(selected_cost),
476 removed_messages,
477 });
478 }
479 if selected_groups.len() == 1 {
480 let available_tokens = budget.saturating_sub(selected_cost);
481 let content = match note {
482 Message::HistoryNote { content } => content,
483 _ => unreachable!(),
484 };
485 let Some(note) =
486 fit_history_note(&content, available_tokens, self.config.bytes_per_token)
487 else {
488 return Err(ContextError(
489 "compaction note cannot fit context budget".into(),
490 ));
491 };
492 let note_cost = note.estimated_tokens(self.config.bytes_per_token);
493 let mut selected: Vec<Message> =
494 selected_groups.into_iter().flatten().collect();
495 selected.insert(0, note);
496 selected_cost += note_cost;
497 return Ok(ContextSelection {
498 messages: selected,
499 before_tokens: static_tokens.saturating_add(before_history_tokens),
500 after_tokens: static_tokens.saturating_add(selected_cost),
501 removed_messages,
502 });
503 }
504 let removed_group = selected_groups.remove(0);
505 let removed_cost: usize = removed_group
506 .iter()
507 .map(|message| message.estimated_tokens(self.config.bytes_per_token))
508 .sum();
509 selected_cost = selected_cost.saturating_sub(removed_cost);
510 removed_messages += removed_group.len();
511 }
512 }
513
514 let selected: Vec<Message> = selected_groups.into_iter().flatten().collect();
515 Ok(ContextSelection {
516 messages: selected,
517 before_tokens: static_tokens.saturating_add(before_history_tokens),
518 after_tokens: static_tokens.saturating_add(selected_cost),
519 removed_messages,
520 })
521 }
522}
523
524#[derive(Debug, Clone)]
525pub struct HistoryLimits {
526 pub max_bytes: usize,
527 pub max_messages: usize,
528 pub note_max_chars: usize,
529}
530
531impl Default for HistoryLimits {
532 fn default() -> Self {
533 Self {
534 max_bytes: 16 * 1024 * 1024,
535 max_messages: 10_000,
536 note_max_chars: 4_000,
537 }
538 }
539}
540
541#[derive(Debug, Clone)]
542pub struct AgentConfig {
543 pub system_prompt: String,
544 pub max_steps: usize,
545 pub history_limits: HistoryLimits,
546}
547
548#[derive(Debug, Clone)]
549pub struct TurnOutcome {
550 pub steps: usize,
551 pub usage: Usage,
552}
553
554#[derive(Debug, Error)]
555pub enum AgentError {
556 #[error("turn cancelled")]
557 Cancelled,
558 #[error("{0}")]
559 Provider(String),
560 #[error("{0}")]
561 ContextLimit(String),
562 #[error("agent reached its maximum step count")]
563 StepLimit,
564 #[error("{0}")]
565 HistoryLimit(String),
566 #[error("{0}")]
567 ResponseLimit(String),
568 #[error("{0}")]
569 ToolLimit(String),
570 #[error("{0}")]
571 Internal(String),
572}
573
574impl AgentError {
575 pub fn code(&self) -> &'static str {
576 match self {
577 Self::Cancelled => "cancelled",
578 Self::Provider(_) => "provider_error",
579 Self::ContextLimit(_) => "context_limit",
580 Self::StepLimit => "step_limit",
581 Self::HistoryLimit(_) => "history_limit",
582 Self::ResponseLimit(_) => "response_limit",
583 Self::ToolLimit(_) => "tool_limit",
584 Self::Internal(_) => "internal_error",
585 }
586 }
587}
588
589pub struct AgentRuntime {
590 provider: Arc<dyn Provider>,
591 tools: Arc<ToolRegistry>,
592 context: Arc<dyn ContextPolicy>,
593 config: AgentConfig,
594 workspace: PathBuf,
595}
596
597impl AgentRuntime {
598 pub fn new(
599 provider: Arc<dyn Provider>,
600 tools: Arc<ToolRegistry>,
601 context: Arc<dyn ContextPolicy>,
602 config: AgentConfig,
603 workspace: PathBuf,
604 ) -> Self {
605 Self {
606 provider,
607 tools,
608 context,
609 config,
610 workspace,
611 }
612 }
613
614 pub fn model(&self) -> &str {
615 self.provider.model()
616 }
617
618 pub async fn run_turn(
619 &self,
620 history: &mut Vec<Message>,
621 prompt: String,
622 sink: Arc<dyn EventSink>,
623 approvals: Arc<dyn ApprovalGate>,
624 cancellation: CancellationToken,
625 ) -> Result<TurnOutcome, AgentError> {
626 let checkpoint = history.clone();
627 let result = self
628 .run_turn_inner(history, prompt, sink, approvals, cancellation)
629 .await;
630 if result.is_err() {
631 *history = checkpoint;
632 }
633 result
634 }
635
636 async fn run_turn_inner(
637 &self,
638 history: &mut Vec<Message>,
639 prompt: String,
640 sink: Arc<dyn EventSink>,
641 approvals: Arc<dyn ApprovalGate>,
642 cancellation: CancellationToken,
643 ) -> Result<TurnOutcome, AgentError> {
644 if cancellation.is_cancelled() {
645 return Err(AgentError::Cancelled);
646 }
647 history.push(Message::User { content: prompt });
648 self.enforce_history_limits(history, sink.as_ref()).await?;
649 let specs = self.tools.specs();
650 let mut usage = Usage::default();
651
652 for step in 1..=self.config.max_steps {
653 if cancellation.is_cancelled() {
654 return Err(AgentError::Cancelled);
655 }
656 let selection = self
657 .context
658 .select(history, &self.config.system_prompt, &specs)
659 .map_err(|error| AgentError::ContextLimit(error.to_string()))?;
660 if selection.removed_messages > 0 {
661 sink.emit(CoreEvent::ContextCompacted {
662 before_tokens: selection.before_tokens,
663 after_tokens: selection.after_tokens,
664 removed_messages: selection.removed_messages,
665 })
666 .await?;
667 }
668 let delta_sink: Arc<dyn TextDeltaSink> = Arc::new(ForwardDeltas {
669 sink: Arc::clone(&sink),
670 });
671 let response = self
672 .provider
673 .complete(
674 ProviderRequest {
675 system_prompt: self.config.system_prompt.clone(),
676 messages: selection.messages,
677 tools: specs.clone(),
678 },
679 delta_sink,
680 cancellation.child_token(),
681 )
682 .await
683 .map_err(map_provider_error)?;
684 usage.add(&response.usage);
685 sink.emit(CoreEvent::AssistantCompleted {
686 content: response.content.clone(),
687 })
688 .await?;
689 let calls = response.tool_calls.clone();
690 history.push(Message::Assistant {
691 content: response.content,
692 tool_calls: response.tool_calls,
693 });
694 self.enforce_history_limits(history, sink.as_ref()).await?;
695 if calls.is_empty() {
696 return Ok(TurnOutcome { steps: step, usage });
697 }
698
699 for call in calls {
700 if cancellation.is_cancelled() {
701 return Err(AgentError::Cancelled);
702 }
703 sink.emit(CoreEvent::ToolProposed {
704 call_id: call.id.clone(),
705 name: call.name.clone(),
706 arguments: call.arguments.clone(),
707 })
708 .await?;
709 let Some(tool) = self.tools.get(&call.name) else {
710 let output = ToolOutput::failure(format!("unknown tool: {}", call.name));
711 sink.emit(CoreEvent::ToolCompleted {
712 call_id: call.id.clone(),
713 name: call.name.clone(),
714 output: output.clone(),
715 })
716 .await?;
717 history.push(Message::Tool {
718 call_id: call.id,
719 name: call.name,
720 content: output.content,
721 is_error: true,
722 });
723 self.enforce_history_limits(history, sink.as_ref()).await?;
724 continue;
725 };
726 let risk = match tool.risk(&call.arguments) {
727 Ok(risk) => risk,
728 Err(error) => {
729 self.record_tool_error(history, sink.as_ref(), &call, error.to_string())
730 .await?;
731 self.enforce_history_limits(history, sink.as_ref()).await?;
732 continue;
733 }
734 };
735 let summary = match tool.approval_summary(&call.arguments) {
736 Ok(summary) => summary,
737 Err(error) => {
738 self.record_tool_error(history, sink.as_ref(), &call, error.to_string())
739 .await?;
740 self.enforce_history_limits(history, sink.as_ref()).await?;
741 continue;
742 }
743 };
744 let approved = approvals
745 .approve(
746 ApprovalRequest {
747 call_id: call.id.clone(),
748 name: call.name.clone(),
749 risk,
750 cwd: self.workspace.clone(),
751 summary,
752 },
753 cancellation.child_token(),
754 )
755 .await?;
756 let output = if approved {
757 sink.emit(CoreEvent::ToolStarted {
758 call_id: call.id.clone(),
759 name: call.name.clone(),
760 })
761 .await?;
762 tool.execute(
763 call.arguments.clone(),
764 ToolContext {
765 workspace: self.workspace.clone(),
766 cancellation: cancellation.child_token(),
767 },
768 )
769 .await
770 .unwrap_or_else(|error| ToolOutput::failure(error.to_string()))
771 } else {
772 ToolOutput::failure("tool call denied by policy or user")
773 };
774 sink.emit(CoreEvent::ToolCompleted {
775 call_id: call.id.clone(),
776 name: call.name.clone(),
777 output: output.clone(),
778 })
779 .await?;
780 history.push(Message::Tool {
781 call_id: call.id,
782 name: call.name,
783 content: output.content,
784 is_error: output.is_error,
785 });
786 self.enforce_history_limits(history, sink.as_ref()).await?;
787 }
788 }
789 Err(AgentError::StepLimit)
790 }
791
792 async fn record_tool_error(
793 &self,
794 history: &mut Vec<Message>,
795 sink: &dyn EventSink,
796 call: &ToolCall,
797 message: String,
798 ) -> Result<(), AgentError> {
799 let output = ToolOutput::failure(message);
800 sink.emit(CoreEvent::ToolCompleted {
801 call_id: call.id.clone(),
802 name: call.name.clone(),
803 output: output.clone(),
804 })
805 .await?;
806 history.push(Message::Tool {
807 call_id: call.id.clone(),
808 name: call.name.clone(),
809 content: output.content,
810 is_error: true,
811 });
812 Ok(())
813 }
814
815 async fn enforce_history_limits(
816 &self,
817 history: &mut Vec<Message>,
818 sink: &dyn EventSink,
819 ) -> Result<(), AgentError> {
820 let limits = &self.config.history_limits;
821 let mut total_removed = 0;
822 while history.len() > limits.max_messages || history_bytes(history) > limits.max_bytes {
823 let latest_user = history
824 .iter()
825 .rposition(|message| matches!(message, Message::User { .. }))
826 .unwrap_or(0);
827 let active = &history[latest_user..];
828 if active.len() > limits.max_messages || history_bytes(active) > limits.max_bytes {
829 return Err(AgentError::HistoryLimit(
830 "active turn exceeds configured session history limit".into(),
831 ));
832 }
833 let first_user = history
834 .iter()
835 .position(|message| matches!(message, Message::User { .. }))
836 .unwrap_or(latest_user);
837 if first_user == latest_user {
838 if matches!(history.first(), Some(Message::HistoryNote { .. })) {
839 history.remove(0);
840 total_removed += 1;
841 continue;
842 }
843 return Err(AgentError::HistoryLimit(
844 "session history cannot be reduced within its configured limit".into(),
845 ));
846 }
847 let end = history[first_user + 1..]
848 .iter()
849 .position(|message| matches!(message, Message::User { .. }))
850 .map(|index| first_user + 1 + index)
851 .ok_or_else(|| {
852 AgentError::HistoryLimit(
853 "session history has no complete group available to trim".into(),
854 )
855 })?;
856 let removed: Vec<Message> = history.drain(..end).collect();
857 total_removed += removed.len();
858 let note = Message::HistoryNote {
859 content: summarize_history_trim(&removed, total_removed, limits.note_max_chars),
860 };
861 if matches!(history.first(), Some(Message::HistoryNote { .. })) {
862 history.remove(0);
863 }
864 history.insert(0, note);
865 }
866 if total_removed > 0 {
867 sink.emit(CoreEvent::SessionTrimmed {
868 removed_messages: total_removed,
869 history_bytes: history_bytes(history),
870 })
871 .await?;
872 }
873 Ok(())
874 }
875}
876
877fn map_provider_error(error: ProviderError) -> AgentError {
878 match error.kind {
879 ProviderErrorKind::Provider => AgentError::Provider(error.message),
880 ProviderErrorKind::ResponseLimit => AgentError::ResponseLimit(error.message),
881 ProviderErrorKind::ToolLimit => AgentError::ToolLimit(error.message),
882 ProviderErrorKind::Cancelled => AgentError::Cancelled,
883 }
884}
885
886struct ForwardDeltas {
887 sink: Arc<dyn EventSink>,
888}
889
890#[async_trait]
891impl TextDeltaSink for ForwardDeltas {
892 async fn push(&self, delta: &str) -> Result<(), ProviderError> {
893 self.sink
894 .emit(CoreEvent::AssistantDelta {
895 content: delta.to_owned(),
896 })
897 .await
898 .map_err(|error| match error {
899 AgentError::Cancelled => {
900 ProviderError::new(ProviderErrorKind::Cancelled, "turn cancelled")
901 }
902 AgentError::ResponseLimit(message) => {
903 ProviderError::new(ProviderErrorKind::ResponseLimit, message)
904 }
905 AgentError::ToolLimit(message) => {
906 ProviderError::new(ProviderErrorKind::ToolLimit, message)
907 }
908 error => ProviderError::new(ProviderErrorKind::Provider, error.to_string()),
909 })
910 }
911}
912
913fn history_bytes(history: &[Message]) -> usize {
914 serde_json::to_vec(history).map_or(usize::MAX, |value| value.len())
915}
916
917fn summarize_history_trim(messages: &[Message], removed: usize, max_chars: usize) -> String {
918 let mut note =
919 format!("[SCV trimmed {removed} earlier canonical messages to enforce session limits.]\n");
920 for message in messages {
921 let (label, content) = match message {
922 Message::User { content } => ("user", content.as_str()),
923 Message::Assistant { content, .. } => ("assistant", content.as_str()),
924 Message::Tool {
925 name,
926 content,
927 is_error,
928 ..
929 } => {
930 let status = if *is_error { "failed" } else { "ok" };
931 note.push_str(&format!("tool {name} ({status}): "));
932 ("", content.as_str())
933 }
934 Message::HistoryNote { content } => ("earlier", content.as_str()),
935 };
936 if !label.is_empty() {
937 note.push_str(label);
938 note.push_str(": ");
939 }
940 note.push_str(&char_tail(content, 160).replace('\n', " "));
941 note.push('\n');
942 if note.chars().count() >= max_chars {
943 break;
944 }
945 }
946 truncate_chars(¬e, max_chars)
947}
948
949fn truncate_chars(value: &str, max_chars: usize) -> String {
950 value.chars().take(max_chars).collect()
951}
952
953fn fit_history_note(
954 content: &str,
955 available_tokens: usize,
956 bytes_per_token: usize,
957) -> Option<Message> {
958 let chars: Vec<char> = content.chars().collect();
959 let mut low = 0usize;
960 let mut high = chars.len();
961 let mut best = None;
962 while low <= high {
963 let middle = low + (high - low) / 2;
964 let candidate = Message::HistoryNote {
965 content: chars[..middle].iter().collect(),
966 };
967 if candidate.estimated_tokens(bytes_per_token) <= available_tokens {
968 best = Some(candidate);
969 low = middle.saturating_add(1);
970 } else if middle == 0 {
971 break;
972 } else {
973 high = middle - 1;
974 }
975 }
976 best
977}
978
979fn char_tail(value: &str, max_chars: usize) -> String {
980 let count = value.chars().count();
981 value
982 .chars()
983 .skip(count.saturating_sub(max_chars))
984 .collect()
985}
986
987#[cfg(test)]
988mod tests {
989 use std::{collections::VecDeque, sync::Mutex};
990
991 use tokio::sync::Notify;
992
993 use super::*;
994
995 struct ScriptedProvider {
996 responses: Mutex<VecDeque<AssistantResponse>>,
997 }
998
999 #[async_trait]
1000 impl Provider for ScriptedProvider {
1001 fn model(&self) -> &str {
1002 "test-model"
1003 }
1004
1005 async fn complete(
1006 &self,
1007 _request: ProviderRequest,
1008 deltas: Arc<dyn TextDeltaSink>,
1009 _cancellation: CancellationToken,
1010 ) -> Result<AssistantResponse, ProviderError> {
1011 let response = self.responses.lock().unwrap().pop_front().unwrap();
1012 deltas.push(&response.content).await?;
1013 Ok(response)
1014 }
1015 }
1016
1017 struct CollectSink(Mutex<Vec<CoreEvent>>);
1018
1019 #[async_trait]
1020 impl EventSink for CollectSink {
1021 async fn emit(&self, event: CoreEvent) -> Result<(), AgentError> {
1022 self.0.lock().unwrap().push(event);
1023 Ok(())
1024 }
1025 }
1026
1027 struct Allow;
1028
1029 #[async_trait]
1030 impl ApprovalGate for Allow {
1031 async fn approve(
1032 &self,
1033 _request: ApprovalRequest,
1034 _cancellation: CancellationToken,
1035 ) -> Result<bool, AgentError> {
1036 Ok(true)
1037 }
1038 }
1039
1040 struct Deny;
1041
1042 #[async_trait]
1043 impl ApprovalGate for Deny {
1044 async fn approve(
1045 &self,
1046 _request: ApprovalRequest,
1047 _cancellation: CancellationToken,
1048 ) -> Result<bool, AgentError> {
1049 Ok(false)
1050 }
1051 }
1052
1053 struct WaitForCancellation {
1054 entered: Arc<Notify>,
1055 }
1056
1057 #[async_trait]
1058 impl ApprovalGate for WaitForCancellation {
1059 async fn approve(
1060 &self,
1061 _request: ApprovalRequest,
1062 cancellation: CancellationToken,
1063 ) -> Result<bool, AgentError> {
1064 self.entered.notify_one();
1065 cancellation.cancelled().await;
1066 Err(AgentError::Cancelled)
1067 }
1068 }
1069
1070 struct EchoTool;
1071
1072 #[async_trait]
1073 impl Tool for EchoTool {
1074 fn spec(&self) -> ToolSpec {
1075 ToolSpec {
1076 name: "echo".into(),
1077 description: "Echo a value".into(),
1078 parameters: serde_json::json!({
1079 "type":"object",
1080 "properties":{"value":{"type":"string"}},
1081 "required":["value"]
1082 }),
1083 }
1084 }
1085
1086 fn risk(&self, arguments: &Value) -> Result<ToolRisk, ToolError> {
1087 arguments
1088 .get("value")
1089 .and_then(Value::as_str)
1090 .ok_or_else(|| ToolError("value must be a string".into()))?;
1091 Ok(ToolRisk::ReadOnly)
1092 }
1093
1094 fn approval_summary(&self, arguments: &Value) -> Result<String, ToolError> {
1095 self.risk(arguments)?;
1096 Ok("Echo a value".into())
1097 }
1098
1099 async fn execute(
1100 &self,
1101 arguments: Value,
1102 _context: ToolContext,
1103 ) -> Result<ToolOutput, ToolError> {
1104 Ok(ToolOutput::success(
1105 arguments["value"].as_str().unwrap_or_default(),
1106 ))
1107 }
1108 }
1109
1110 #[tokio::test]
1111 async fn completes_a_simple_turn() {
1112 let provider = Arc::new(ScriptedProvider {
1113 responses: Mutex::new(VecDeque::from([AssistantResponse {
1114 content: "done".into(),
1115 tool_calls: Vec::new(),
1116 usage: Usage {
1117 input_tokens: Some(3),
1118 output_tokens: Some(1),
1119 },
1120 }])),
1121 });
1122 let runtime = AgentRuntime::new(
1123 provider,
1124 Arc::new(ToolRegistry::default()),
1125 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1126 AgentConfig {
1127 system_prompt: "test".into(),
1128 max_steps: 2,
1129 history_limits: HistoryLimits::default(),
1130 },
1131 PathBuf::from("/tmp"),
1132 );
1133 let sink = Arc::new(CollectSink(Mutex::new(Vec::new())));
1134 let mut history = Vec::new();
1135 let outcome = runtime
1136 .run_turn(
1137 &mut history,
1138 "hello".into(),
1139 sink.clone(),
1140 Arc::new(Allow),
1141 CancellationToken::new(),
1142 )
1143 .await
1144 .unwrap();
1145 assert_eq!(outcome.steps, 1);
1146 assert_eq!(history.len(), 2);
1147 assert!(matches!(
1148 sink.0.lock().unwrap().last(),
1149 Some(CoreEvent::AssistantCompleted { .. })
1150 ));
1151 }
1152
1153 #[test]
1154 fn context_keeps_tool_groups_together() {
1155 let policy = BudgetContextPolicy::new(ContextConfig {
1156 max_tokens: 120,
1157 reserve_output_tokens: 10,
1158 safety_margin_tokens: 10,
1159 bytes_per_token: 3,
1160 summary_max_chars: 120,
1161 })
1162 .unwrap();
1163 let history = vec![
1164 Message::User {
1165 content: "old request ".repeat(20),
1166 },
1167 Message::Assistant {
1168 content: String::new(),
1169 tool_calls: vec![ToolCall {
1170 id: "1".into(),
1171 name: "read".into(),
1172 arguments: serde_json::json!({"path":"a"}),
1173 }],
1174 },
1175 Message::Tool {
1176 call_id: "1".into(),
1177 name: "read".into(),
1178 content: "result".into(),
1179 is_error: false,
1180 },
1181 Message::User {
1182 content: "new".into(),
1183 },
1184 ];
1185 let selection = policy.select(&history, "system", &[]).unwrap();
1186 assert!(selection.removed_messages > 0);
1187 assert_eq!(
1188 selection.removed_messages,
1189 history.len() - (selection.messages.len() - 1)
1190 );
1191 assert!(matches!(
1192 selection.messages.last(),
1193 Some(Message::User { .. })
1194 ));
1195 assert!(
1196 !selection
1197 .messages
1198 .iter()
1199 .any(|message| matches!(message, Message::Tool { call_id, .. } if call_id == "1"))
1200 );
1201 }
1202
1203 #[tokio::test]
1204 async fn repeated_history_trimming_rebuilds_the_note_and_makes_progress() {
1205 let runtime = AgentRuntime::new(
1206 Arc::new(ScriptedProvider {
1207 responses: Mutex::new(VecDeque::new()),
1208 }),
1209 Arc::new(ToolRegistry::default()),
1210 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1211 AgentConfig {
1212 system_prompt: "test".into(),
1213 max_steps: 1,
1214 history_limits: HistoryLimits {
1215 max_bytes: 4096,
1216 max_messages: 3,
1217 note_max_chars: 80,
1218 },
1219 },
1220 PathBuf::from("/tmp"),
1221 );
1222 let sink = CollectSink(Mutex::new(Vec::new()));
1223 let mut history = vec![
1224 Message::HistoryNote {
1225 content: "previous trim".into(),
1226 },
1227 Message::User {
1228 content: "old request".into(),
1229 },
1230 Message::Assistant {
1231 content: "old answer".into(),
1232 tool_calls: Vec::new(),
1233 },
1234 Message::User {
1235 content: "active request".into(),
1236 },
1237 ];
1238 runtime
1239 .enforce_history_limits(&mut history, &sink)
1240 .await
1241 .unwrap();
1242 assert!(history.len() <= 3);
1243 assert!(matches!(history.first(), Some(Message::HistoryNote { .. })));
1244 assert!(matches!(history.last(), Some(Message::User { .. })));
1245 }
1246
1247 #[tokio::test]
1248 async fn active_turn_over_history_limit_rolls_back() {
1249 let runtime = AgentRuntime::new(
1250 Arc::new(ScriptedProvider {
1251 responses: Mutex::new(VecDeque::new()),
1252 }),
1253 Arc::new(ToolRegistry::default()),
1254 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1255 AgentConfig {
1256 system_prompt: "test".into(),
1257 max_steps: 1,
1258 history_limits: HistoryLimits {
1259 max_bytes: 16,
1260 max_messages: 10,
1261 note_max_chars: 8,
1262 },
1263 },
1264 PathBuf::from("/tmp"),
1265 );
1266 let sink = CollectSink(Mutex::new(Vec::new()));
1267 let mut history = Vec::new();
1268 let result = runtime
1269 .run_turn(
1270 &mut history,
1271 "too large for the configured history".into(),
1272 Arc::new(sink),
1273 Arc::new(Allow),
1274 CancellationToken::new(),
1275 )
1276 .await;
1277 assert!(matches!(result, Err(AgentError::HistoryLimit(_))));
1278 assert!(history.is_empty());
1279 }
1280
1281 #[tokio::test]
1282 async fn executes_a_multi_step_tool_loop_and_aggregates_usage() {
1283 let provider = Arc::new(ScriptedProvider {
1284 responses: Mutex::new(VecDeque::from([
1285 AssistantResponse {
1286 content: String::new(),
1287 tool_calls: vec![ToolCall {
1288 id: "call-1".into(),
1289 name: "echo".into(),
1290 arguments: serde_json::json!({"value":"hello"}),
1291 }],
1292 usage: Usage {
1293 input_tokens: Some(2),
1294 output_tokens: Some(1),
1295 },
1296 },
1297 AssistantResponse {
1298 content: "done".into(),
1299 tool_calls: Vec::new(),
1300 usage: Usage {
1301 input_tokens: Some(4),
1302 output_tokens: Some(2),
1303 },
1304 },
1305 ])),
1306 });
1307 let mut registry = ToolRegistry::default();
1308 registry.register(Arc::new(EchoTool)).unwrap();
1309 let runtime = AgentRuntime::new(
1310 provider,
1311 Arc::new(registry),
1312 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1313 AgentConfig {
1314 system_prompt: "test".into(),
1315 max_steps: 3,
1316 history_limits: HistoryLimits::default(),
1317 },
1318 PathBuf::from("/tmp"),
1319 );
1320 let sink = Arc::new(CollectSink(Mutex::new(Vec::new())));
1321 let mut history = Vec::new();
1322 let outcome = runtime
1323 .run_turn(
1324 &mut history,
1325 "start".into(),
1326 sink,
1327 Arc::new(Allow),
1328 CancellationToken::new(),
1329 )
1330 .await
1331 .unwrap();
1332 assert_eq!(outcome.steps, 2);
1333 assert_eq!(outcome.usage.input_tokens, Some(6));
1334 assert_eq!(outcome.usage.output_tokens, Some(3));
1335 assert!(matches!(
1336 history.get(2),
1337 Some(Message::Tool {
1338 content,
1339 is_error: false,
1340 ..
1341 }) if content == "hello"
1342 ));
1343 }
1344
1345 #[tokio::test]
1346 async fn denial_is_recorded_as_a_model_visible_tool_failure() {
1347 let provider = Arc::new(ScriptedProvider {
1348 responses: Mutex::new(VecDeque::from([
1349 AssistantResponse {
1350 content: String::new(),
1351 tool_calls: vec![ToolCall {
1352 id: "call-1".into(),
1353 name: "echo".into(),
1354 arguments: serde_json::json!({"value":"blocked"}),
1355 }],
1356 usage: Usage::default(),
1357 },
1358 AssistantResponse {
1359 content: "handled".into(),
1360 tool_calls: Vec::new(),
1361 usage: Usage::default(),
1362 },
1363 ])),
1364 });
1365 let mut registry = ToolRegistry::default();
1366 registry.register(Arc::new(EchoTool)).unwrap();
1367 let runtime = AgentRuntime::new(
1368 provider,
1369 Arc::new(registry),
1370 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1371 AgentConfig {
1372 system_prompt: "test".into(),
1373 max_steps: 3,
1374 history_limits: HistoryLimits::default(),
1375 },
1376 PathBuf::from("/tmp"),
1377 );
1378 let mut history = Vec::new();
1379 runtime
1380 .run_turn(
1381 &mut history,
1382 "start".into(),
1383 Arc::new(CollectSink(Mutex::new(Vec::new()))),
1384 Arc::new(Deny),
1385 CancellationToken::new(),
1386 )
1387 .await
1388 .unwrap();
1389 assert!(matches!(
1390 history.get(2),
1391 Some(Message::Tool {
1392 content,
1393 is_error: true,
1394 ..
1395 }) if content.contains("denied")
1396 ));
1397 }
1398
1399 #[tokio::test]
1400 async fn cancellation_during_approval_rolls_back_the_active_tool_group() {
1401 let provider = Arc::new(ScriptedProvider {
1402 responses: Mutex::new(VecDeque::from([AssistantResponse {
1403 content: String::new(),
1404 tool_calls: vec![ToolCall {
1405 id: "call-cancel".into(),
1406 name: "echo".into(),
1407 arguments: serde_json::json!({"value":"hello"}),
1408 }],
1409 usage: Usage::default(),
1410 }])),
1411 });
1412 let mut registry = ToolRegistry::default();
1413 registry.register(Arc::new(EchoTool)).unwrap();
1414 let runtime = AgentRuntime::new(
1415 provider,
1416 Arc::new(registry),
1417 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1418 AgentConfig {
1419 system_prompt: "test".into(),
1420 max_steps: 2,
1421 history_limits: HistoryLimits::default(),
1422 },
1423 PathBuf::from("/tmp"),
1424 );
1425 let before = vec![
1426 Message::User {
1427 content: "previous".into(),
1428 },
1429 Message::Assistant {
1430 content: "answer".into(),
1431 tool_calls: Vec::new(),
1432 },
1433 ];
1434 let mut history = before.clone();
1435 let cancellation = CancellationToken::new();
1436 let cancel = cancellation.clone();
1437 let entered = Arc::new(Notify::new());
1438 let wait = Arc::clone(&entered);
1439 let run = runtime.run_turn(
1440 &mut history,
1441 "new turn".into(),
1442 Arc::new(CollectSink(Mutex::new(Vec::new()))),
1443 Arc::new(WaitForCancellation { entered }),
1444 cancellation,
1445 );
1446 let cancel_when_waiting = async move {
1447 wait.notified().await;
1448 cancel.cancel();
1449 };
1450 let (result, ()) = tokio::join!(run, cancel_when_waiting);
1451 assert!(matches!(result, Err(AgentError::Cancelled)));
1452 assert_eq!(history, before);
1453 }
1454
1455 #[tokio::test]
1456 async fn stops_after_the_configured_maximum_step() {
1457 let provider = Arc::new(ScriptedProvider {
1458 responses: Mutex::new(VecDeque::from([AssistantResponse {
1459 content: String::new(),
1460 tool_calls: vec![ToolCall {
1461 id: "call-1".into(),
1462 name: "echo".into(),
1463 arguments: serde_json::json!({"value":"one"}),
1464 }],
1465 usage: Usage::default(),
1466 }])),
1467 });
1468 let mut registry = ToolRegistry::default();
1469 registry.register(Arc::new(EchoTool)).unwrap();
1470 let runtime = AgentRuntime::new(
1471 provider,
1472 Arc::new(registry),
1473 Arc::new(BudgetContextPolicy::new(ContextConfig::default()).unwrap()),
1474 AgentConfig {
1475 system_prompt: "test".into(),
1476 max_steps: 1,
1477 history_limits: HistoryLimits::default(),
1478 },
1479 PathBuf::from("/tmp"),
1480 );
1481 let result = runtime
1482 .run_turn(
1483 &mut Vec::new(),
1484 "start".into(),
1485 Arc::new(CollectSink(Mutex::new(Vec::new()))),
1486 Arc::new(Allow),
1487 CancellationToken::new(),
1488 )
1489 .await;
1490 assert!(matches!(result, Err(AgentError::StepLimit)));
1491 }
1492
1493 #[test]
1494 fn duplicate_tool_registration_does_not_replace_the_original() {
1495 let mut registry = ToolRegistry::default();
1496 registry.register(Arc::new(EchoTool)).unwrap();
1497 assert!(registry.register(Arc::new(EchoTool)).is_err());
1498 assert_eq!(registry.tools.len(), 1);
1499 assert!(registry.get("echo").is_some());
1500 }
1501}