1use crate::auth::AuthStorage;
16use crate::compaction::{self, ResolvedCompactionSettings};
17use crate::compaction_worker::{
18 CompactionAdmissionSignals, CompactionQuota, CompactionWorkerState,
19};
20use crate::error::{Error, Result};
21use crate::extension_events::{
22 BeforeAgentStartOutcome, InputEventOutcome, SessionBeforeCompactOutcome,
23 apply_before_agent_start_response, apply_input_event_response,
24 apply_session_before_compact_response,
25};
26use crate::extension_tools::collect_extension_tool_wrappers;
27use crate::extensions::{
28 EXTENSION_EVENT_TIMEOUT_MS, ExtensionAiCompletionRequest, ExtensionDeliverAs,
29 ExtensionEventName, ExtensionHostActions, ExtensionLoadSpec, ExtensionManager, ExtensionPolicy,
30 ExtensionRegion, ExtensionRuntimeHandle, ExtensionSendMessage, ExtensionSendUserMessage,
31 JsExtensionLoadSpec, JsExtensionRuntimeHandle, NativeRustExtensionLoadSpec,
32 NativeRustExtensionRuntimeHandle, RepairPolicyMode, resolve_extension_load_spec,
33};
34#[cfg(feature = "wasm-host")]
35use crate::extensions::{WasmExtensionHost, WasmExtensionLoadSpec};
36use crate::extensions_js::{PiJsRuntimeConfig, RepairMode};
37use crate::model::{
38 AssistantMessage, AssistantMessageEvent, ContentBlock, CustomMessage, ImageContent, Message,
39 StopReason, StreamEvent, TextContent, ThinkingContent, ToolCall, ToolResultMessage, Usage,
40 UserContent, UserMessage,
41};
42use crate::models::{
43 ModelEntry, ModelRegistry, model_requires_configured_credential, normalize_api_key_opt,
44};
45use crate::provider::{Context, Provider, StreamOptions, ToolDef};
46use crate::semantic_workspace_graph::{ContextBundleItem, SemanticContextBundle};
47use crate::session::{AutosaveFlushTrigger, Session, SessionHandle};
48use crate::tools::{Tool, ToolEffects, ToolOutput, ToolRegistry, ToolUpdate};
49use asupersync::runtime::{Runtime, RuntimeBuilder, RuntimeHandle};
50use asupersync::sync::{Mutex, Notify, OwnedMutexGuard};
51use async_trait::async_trait;
52use chrono::Utc;
53use futures::FutureExt;
54use futures::StreamExt;
55use futures::future::BoxFuture;
56use futures::stream;
57use serde::Serialize;
58use serde_json::{Value, json};
59use sha2::{Digest as _, Sha256};
60use std::borrow::Cow;
61use std::collections::VecDeque;
62use std::fmt;
63use std::sync::Arc;
64use std::sync::Mutex as StdMutex;
65use std::sync::OnceLock;
66use std::sync::atomic::{AtomicBool, Ordering};
67use std::time::{Duration, Instant};
68use tracing::warn;
69
70const MIN_COMPATIBLE_TOOL_PARALLELISM: usize = 8;
71const MAX_AUTO_COMPATIBLE_TOOL_PARALLELISM: usize = 64;
72const MAX_CONFIGURED_COMPATIBLE_TOOL_PARALLELISM: usize = 256;
73const MAX_STEERING_QUEUE_SIZE: usize = 100;
75const MAX_FOLLOW_UP_QUEUE_SIZE: usize = 100;
77const MAX_AGENT_MESSAGES: usize = 10_000;
79pub const TURN_LATENCY_BREAKDOWN_SCHEMA_V1: &str = "pi.agent.turn_latency_breakdown.v1";
81pub const TOOL_EFFECT_BATCH_PLAN_SCHEMA_V1: &str = "pi.agent.tool_effect_batch_plan.v1";
83const TOOL_CANCELLATION_SCHEMA_V1: &str = "pi.tool.cancellation.v1";
84const TOOL_APPROVAL_DENIED_SCHEMA_V1: &str = "pi.tool.approval_denied.v1";
85const TOOL_APPROVAL_STATUS_SCHEMA_V1: &str = "pi.tool.approval_status.v1";
86const SEMANTIC_CONTEXT_PROMPT_SCHEMA_V1: &str = "pi.semantic_context_prompt.v1";
87const SEMANTIC_CONTEXT_PROVENANCE_SCHEMA_V1: &str = "pi.semantic_context_provenance.v1";
88const SEMANTIC_CONTEXT_CUSTOM_TYPE: &str = "semantic_context_bundle";
89const DEFAULT_SEMANTIC_CONTEXT_PROMPT_MAX_BYTES: u64 = 16 * 1024;
90const DEFAULT_SEMANTIC_CONTEXT_PROMPT_MAX_ITEMS: usize = 16;
91
92fn compatible_tool_parallelism_limit() -> usize {
93 static LIMIT: OnceLock<usize> = OnceLock::new();
94 *LIMIT.get_or_init(|| {
95 let host_parallelism = std::thread::available_parallelism()
96 .map_or(MIN_COMPATIBLE_TOOL_PARALLELISM, |parallelism| {
97 parallelism.get()
98 });
99 resolve_compatible_tool_parallelism(
100 std::env::var("PI_MAX_CONCURRENT_COMPATIBLE_TOOLS")
101 .ok()
102 .as_deref(),
103 host_parallelism,
104 )
105 })
106}
107
108fn resolve_compatible_tool_parallelism(
109 raw_override: Option<&str>,
110 host_parallelism: usize,
111) -> usize {
112 let host_default = host_parallelism.clamp(
113 MIN_COMPATIBLE_TOOL_PARALLELISM,
114 MAX_AUTO_COMPATIBLE_TOOL_PARALLELISM,
115 );
116
117 let Some(raw) = raw_override.map(str::trim).filter(|raw| !raw.is_empty()) else {
118 return host_default;
119 };
120
121 match raw.parse::<usize>() {
122 Ok(0) => {
123 warn!(
124 value = raw,
125 "Ignoring PI_MAX_CONCURRENT_COMPATIBLE_TOOLS=0; using host-scaled default"
126 );
127 host_default
128 }
129 Ok(limit) => limit.clamp(1, MAX_CONFIGURED_COMPATIBLE_TOOL_PARALLELISM),
130 Err(err) => {
131 warn!(
132 value = raw,
133 error = %err,
134 "Ignoring invalid PI_MAX_CONCURRENT_COMPATIBLE_TOOLS; using host-scaled default"
135 );
136 host_default
137 }
138 }
139}
140
141fn duration_millis_saturating(duration: Duration) -> u64 {
142 u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
143}
144
145fn duration_micros_saturating(duration: Duration) -> u64 {
146 u64::try_from(duration.as_micros()).unwrap_or(u64::MAX)
147}
148
149fn record_global_latency(counter: &crate::session_metrics::TimingCounter, duration: Duration) {
150 if crate::session_metrics::global().enabled() {
151 counter.record(duration_micros_saturating(duration));
152 }
153}
154
155#[derive(Debug, Clone, Serialize, Default)]
157#[serde(rename_all = "camelCase")]
158pub struct LatencyPercentiles {
159 pub p50_ms: u64,
161 pub p95_ms: u64,
163 pub p99_ms: u64,
165 pub p999_ms: u64,
167}
168
169impl LatencyPercentiles {
170 fn from_samples(samples: &[u64]) -> Self {
171 Self {
172 p50_ms: percentile_nearest_rank(samples, 50),
173 p95_ms: percentile_nearest_rank(samples, 95),
174 p99_ms: percentile_nearest_rank(samples, 99),
175 p999_ms: percentile_nearest_rank_per_mille(samples, 999),
176 }
177 }
178}
179
180#[derive(Debug, Clone, Serialize, Default)]
182#[serde(rename_all = "camelCase")]
183pub struct LatencyComponentBreakdown {
184 pub duration_ms: u64,
186 pub samples: usize,
188 pub tail_percentiles: LatencyPercentiles,
190}
191
192impl LatencyComponentBreakdown {
193 #[must_use]
195 pub fn from_millis_samples(samples: &[u64]) -> Self {
196 Self {
197 duration_ms: samples.iter().copied().fold(0u64, u64::saturating_add),
198 samples: samples.len(),
199 tail_percentiles: LatencyPercentiles::from_samples(samples),
200 }
201 }
202}
203
204#[derive(Debug, Clone, Serialize)]
206#[serde(rename_all = "camelCase")]
207pub struct TurnLatencyBreakdown {
208 pub schema: &'static str,
210 pub total_ms: u64,
212 pub provider_streaming: LatencyComponentBreakdown,
214 pub local_tools: LatencyComponentBreakdown,
216 pub extension_hostcalls: LatencyComponentBreakdown,
218 pub persistence: LatencyComponentBreakdown,
220 pub dominant_component: String,
222}
223
224impl TurnLatencyBreakdown {
225 #[must_use]
227 pub fn from_component_samples(
228 total_ms: u64,
229 provider_streaming_ms: &[u64],
230 local_tool_ms: &[u64],
231 extension_hostcall_ms: &[u64],
232 persistence_ms: &[u64],
233 ) -> Self {
234 let provider_streaming =
235 LatencyComponentBreakdown::from_millis_samples(provider_streaming_ms);
236 let local_tools = LatencyComponentBreakdown::from_millis_samples(local_tool_ms);
237 let extension_hostcalls =
238 LatencyComponentBreakdown::from_millis_samples(extension_hostcall_ms);
239 let persistence = LatencyComponentBreakdown::from_millis_samples(persistence_ms);
240 let dominant_component = dominant_latency_component(
241 &provider_streaming,
242 &local_tools,
243 &extension_hostcalls,
244 &persistence,
245 );
246
247 Self {
248 schema: TURN_LATENCY_BREAKDOWN_SCHEMA_V1,
249 total_ms,
250 provider_streaming,
251 local_tools,
252 extension_hostcalls,
253 persistence,
254 dominant_component,
255 }
256 }
257}
258
259fn percentile_nearest_rank(samples: &[u64], percentile: usize) -> u64 {
260 if samples.is_empty() {
261 return 0;
262 }
263
264 let mut sorted = samples.to_vec();
265 sorted.sort_unstable();
266 let len = sorted.len();
267 let rank = percentile
268 .saturating_mul(len)
269 .div_ceil(100)
270 .saturating_sub(1)
271 .min(len.saturating_sub(1));
272 sorted[rank]
273}
274
275fn percentile_nearest_rank_per_mille(samples: &[u64], permille: usize) -> u64 {
276 if samples.is_empty() {
277 return 0;
278 }
279
280 let mut sorted = samples.to_vec();
281 sorted.sort_unstable();
282 let len = sorted.len();
283 let rank = permille
284 .saturating_mul(len)
285 .div_ceil(1000)
286 .saturating_sub(1)
287 .min(len.saturating_sub(1));
288 sorted[rank]
289}
290
291fn dominant_latency_component(
292 provider_streaming: &LatencyComponentBreakdown,
293 local_tools: &LatencyComponentBreakdown,
294 extension_hostcalls: &LatencyComponentBreakdown,
295 persistence: &LatencyComponentBreakdown,
296) -> String {
297 [
298 ("provider_streaming", provider_streaming.duration_ms),
299 ("local_tools", local_tools.duration_ms),
300 ("extension_hostcalls", extension_hostcalls.duration_ms),
301 ("persistence", persistence.duration_ms),
302 ]
303 .into_iter()
304 .max_by_key(|(_, duration_ms)| *duration_ms)
305 .filter(|(_, duration_ms)| *duration_ms > 0)
306 .map_or_else(|| "none".to_string(), |(name, _)| name.to_string())
307}
308
309#[derive(Debug)]
310struct TurnLatencyAccumulator {
311 started_at: Instant,
312 provider_streaming_ms: Vec<u64>,
313 local_tool_ms: Vec<u64>,
314 extension_hostcall_ms: Vec<u64>,
315 persistence_ms: Vec<u64>,
316}
317
318impl TurnLatencyAccumulator {
319 fn started() -> Self {
320 Self {
321 started_at: Instant::now(),
322 provider_streaming_ms: Vec::new(),
323 local_tool_ms: Vec::new(),
324 extension_hostcall_ms: Vec::new(),
325 persistence_ms: Vec::new(),
326 }
327 }
328
329 fn snapshot(&self) -> TurnLatencyBreakdown {
330 TurnLatencyBreakdown::from_component_samples(
331 duration_millis_saturating(self.started_at.elapsed()),
332 &self.provider_streaming_ms,
333 &self.local_tool_ms,
334 &self.extension_hostcall_ms,
335 &self.persistence_ms,
336 )
337 }
338}
339
340type SharedTurnLatencyAccumulator = Arc<StdMutex<TurnLatencyAccumulator>>;
341
342fn snapshot_turn_latency(
343 latency: &SharedTurnLatencyAccumulator,
344) -> Option<Box<TurnLatencyBreakdown>> {
345 latency.lock().ok().map(|guard| Box::new(guard.snapshot()))
346}
347
348fn record_provider_streaming_latency(latency: &SharedTurnLatencyAccumulator, duration: Duration) {
349 if let Ok(mut guard) = latency.lock() {
350 guard
351 .provider_streaming_ms
352 .push(duration_millis_saturating(duration));
353 }
354 let metrics = crate::session_metrics::global();
355 record_global_latency(&metrics.provider_streaming, duration);
356}
357
358fn record_local_tool_latency(latency: &SharedTurnLatencyAccumulator, duration: Duration) {
359 if let Ok(mut guard) = latency.lock() {
360 guard
361 .local_tool_ms
362 .push(duration_millis_saturating(duration));
363 }
364 let metrics = crate::session_metrics::global();
365 record_global_latency(&metrics.local_tools, duration);
366}
367
368fn record_extension_hostcall_latency(latency: &SharedTurnLatencyAccumulator, duration: Duration) {
369 if let Ok(mut guard) = latency.lock() {
370 guard
371 .extension_hostcall_ms
372 .push(duration_millis_saturating(duration));
373 }
374 let metrics = crate::session_metrics::global();
375 record_global_latency(&metrics.extension_hostcalls, duration);
376}
377
378#[derive(Debug, Clone, Copy, PartialEq, Eq)]
379struct ToolEffectBatch {
380 start: usize,
381 end: usize,
382}
383
384#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
386#[serde(rename_all = "camelCase")]
387pub struct ToolEffectBatchEvidence {
388 pub start: usize,
390 pub end: usize,
392 pub len: usize,
394 pub combined_effects: Vec<&'static str>,
396 pub parallel_safe: bool,
398 #[serde(skip_serializing_if = "Option::is_none")]
400 pub barrier_reason: Option<&'static str>,
401}
402
403#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
405#[serde(rename_all = "camelCase")]
406pub struct ToolEffectBatchPlanEvidence {
407 pub schema: &'static str,
409 pub tool_count: usize,
411 pub parallelism_cap: usize,
413 pub batches: Vec<ToolEffectBatchEvidence>,
415}
416
417fn plan_tool_effect_batches(effects: &[ToolEffects]) -> Vec<ToolEffectBatch> {
418 let Some((&first_effects, remaining_effects)) = effects.split_first() else {
419 return Vec::new();
420 };
421
422 let mut batches = Vec::new();
423 let mut start = 0;
424 let mut active_effects = first_effects;
425
426 for (offset, candidate_effects) in remaining_effects.iter().copied().enumerate() {
427 let index = offset + 1;
428 if active_effects.compatible_with(candidate_effects) {
429 active_effects = active_effects.union(candidate_effects);
430 } else {
431 batches.push(ToolEffectBatch { start, end: index });
432 start = index;
433 active_effects = candidate_effects;
434 }
435 }
436
437 batches.push(ToolEffectBatch {
438 start,
439 end: effects.len(),
440 });
441 batches
442}
443
444fn combined_tool_effects(effects: &[ToolEffects]) -> Option<ToolEffects> {
445 effects.iter().copied().reduce(ToolEffects::union)
446}
447
448const fn tool_effect_barrier_reason(effects: ToolEffects) -> Option<&'static str> {
449 if effects.parallel_safe() {
450 return None;
451 }
452 match (effects.writes(), effects.appends(), effects.processes()) {
453 (true, true, true) => Some("write_append_process_barrier"),
454 (true, true, false) => Some("write_append_barrier"),
455 (true, false, true) => Some("write_process_barrier"),
456 (false, true, true) => Some("append_process_barrier"),
457 (true, false, false) => Some("write_barrier"),
458 (false, true, false) => Some("append_barrier"),
459 (false, false, true) => Some("process_barrier"),
460 (false, false, false) => Some("undeclared_effects_barrier"),
461 }
462}
463
464#[must_use]
466pub fn tool_effect_batch_plan_evidence(
467 effects: &[ToolEffects],
468 parallelism_cap: usize,
469) -> ToolEffectBatchPlanEvidence {
470 let batches = plan_tool_effect_batches(effects)
471 .into_iter()
472 .map(|batch| {
473 let combined_effects = effects
474 .get(batch.start..batch.end)
475 .and_then(combined_tool_effects)
476 .unwrap_or_else(ToolEffects::read);
477 ToolEffectBatchEvidence {
478 start: batch.start,
479 end: batch.end,
480 len: batch.end.saturating_sub(batch.start),
481 combined_effects: combined_effects.labels(),
482 parallel_safe: combined_effects.parallel_safe(),
483 barrier_reason: tool_effect_barrier_reason(combined_effects),
484 }
485 })
486 .collect();
487
488 ToolEffectBatchPlanEvidence {
489 schema: TOOL_EFFECT_BATCH_PLAN_SCHEMA_V1,
490 tool_count: effects.len(),
491 parallelism_cap,
492 batches,
493 }
494}
495
496pub const MAX_TOOL_ITERATIONS_DEFAULT: usize = 50;
508
509pub const MAX_TOOL_ITERATIONS_CEILING: usize = 1_000;
515
516const ITERATION_WARN_NUMERATOR: usize = 4;
521const ITERATION_WARN_DENOMINATOR: usize = 5;
522
523const ITERATION_WARN_MIN_CAP: usize = 5;
527
528pub fn resolved_max_tool_iterations_default() -> usize {
534 resolve_max_tool_iterations(std::env::var("PI_MAX_TOOL_ITERATIONS").ok().as_deref())
535}
536
537pub fn resolve_max_tool_iterations(raw_override: Option<&str>) -> usize {
543 let Some(raw) = raw_override.map(str::trim).filter(|raw| !raw.is_empty()) else {
544 return MAX_TOOL_ITERATIONS_DEFAULT;
545 };
546 match raw.parse::<usize>() {
547 Ok(0) => {
548 warn!(
549 "PI_MAX_TOOL_ITERATIONS=0 is invalid; falling back to {}",
550 MAX_TOOL_ITERATIONS_DEFAULT
551 );
552 MAX_TOOL_ITERATIONS_DEFAULT
553 }
554 Ok(n) if n > MAX_TOOL_ITERATIONS_CEILING => {
555 warn!(
556 "PI_MAX_TOOL_ITERATIONS={n} exceeds ceiling {MAX_TOOL_ITERATIONS_CEILING}; clamping to {MAX_TOOL_ITERATIONS_CEILING}"
557 );
558 MAX_TOOL_ITERATIONS_CEILING
559 }
560 Ok(n) => n,
561 Err(err) => {
562 warn!(
563 "PI_MAX_TOOL_ITERATIONS={raw:?} is not a valid usize ({err}); falling back to {}",
564 MAX_TOOL_ITERATIONS_DEFAULT
565 );
566 MAX_TOOL_ITERATIONS_DEFAULT
567 }
568 }
569}
570
571pub fn clamp_max_tool_iterations(value: Option<usize>) -> usize {
578 match value {
579 None => MAX_TOOL_ITERATIONS_DEFAULT,
580 Some(0) => {
581 warn!(
582 "--max-tool-iterations=0 is invalid; falling back to {}",
583 MAX_TOOL_ITERATIONS_DEFAULT
584 );
585 MAX_TOOL_ITERATIONS_DEFAULT
586 }
587 Some(n) if n > MAX_TOOL_ITERATIONS_CEILING => {
588 warn!(
589 "--max-tool-iterations={n} exceeds ceiling {MAX_TOOL_ITERATIONS_CEILING}; clamping to {MAX_TOOL_ITERATIONS_CEILING}"
590 );
591 MAX_TOOL_ITERATIONS_CEILING
592 }
593 Some(n) => n,
594 }
595}
596
597pub const fn should_warn_at_iteration_threshold(current: usize, max: usize) -> bool {
608 max >= ITERATION_WARN_MIN_CAP
609 && current >= max.saturating_mul(ITERATION_WARN_NUMERATOR) / ITERATION_WARN_DENOMINATOR
610}
611
612pub fn iteration_handoff_steering_text(current: usize, max: usize) -> String {
616 format!(
617 "[runtime] Tool-iteration budget at >=80% (used {current} of {max}). \
618 Per the iteration-aware-handoff protocol in your spec, begin graceful \
619 handoff now: commit current work, post a one-line status note, and \
620 write an incomplete-handoff envelope with what's done / what remains \
621 / next-agent starting position. Do NOT compress remaining work into \
622 the last few iterations."
623 )
624}
625
626#[derive(Clone)]
628pub struct AgentConfig {
629 pub system_prompt: Option<String>,
631
632 pub max_tool_iterations: usize,
634
635 pub stream_options: StreamOptions,
637
638 pub block_images: bool,
640
641 pub fail_closed_hooks: bool,
643
644 pub tool_approval: Option<ToolApprovalHandler>,
646}
647
648impl fmt::Debug for AgentConfig {
649 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
650 f.debug_struct("AgentConfig")
651 .field("system_prompt", &self.system_prompt)
652 .field("max_tool_iterations", &self.max_tool_iterations)
653 .field("stream_options", &self.stream_options)
654 .field("block_images", &self.block_images)
655 .field("fail_closed_hooks", &self.fail_closed_hooks)
656 .field("tool_approval", &self.tool_approval.is_some())
657 .finish()
658 }
659}
660
661#[derive(Debug, Clone)]
663pub struct ToolApprovalRequest {
664 pub tool_call_id: String,
665 pub tool_name: String,
666 pub arguments: Value,
667}
668
669#[derive(Debug, Clone, PartialEq, Eq)]
671pub enum ToolApprovalDecision {
672 Allow,
673 Deny { reason: String },
674}
675
676impl ToolApprovalDecision {
677 #[must_use]
678 pub fn deny(reason: impl Into<String>) -> Self {
679 Self::Deny {
680 reason: reason.into(),
681 }
682 }
683}
684
685pub type ToolApprovalHandler =
686 Arc<dyn Fn(ToolApprovalRequest) -> BoxFuture<'static, ToolApprovalDecision> + Send + Sync>;
687
688impl Default for AgentConfig {
689 fn default() -> Self {
690 Self {
691 system_prompt: None,
692 max_tool_iterations: resolved_max_tool_iterations_default(),
693 stream_options: StreamOptions::default(),
694 block_images: false,
695 fail_closed_hooks: false,
696 tool_approval: None,
697 }
698 }
699}
700
701#[derive(Debug, Clone)]
703pub struct SemanticContextBundleInjection {
704 pub enabled: bool,
705 pub bundle: SemanticContextBundle,
706 pub max_prompt_items: usize,
707 pub max_prompt_bytes: u64,
708 pub include_exclusion_summary: bool,
709 pub include_validation_commands: bool,
710}
711
712impl SemanticContextBundleInjection {
713 pub fn enabled(bundle: SemanticContextBundle) -> Self {
714 let max_prompt_items = bundle
715 .budget
716 .max_items
717 .min(DEFAULT_SEMANTIC_CONTEXT_PROMPT_MAX_ITEMS);
718 let max_prompt_bytes = bundle
719 .budget
720 .max_bytes
721 .min(DEFAULT_SEMANTIC_CONTEXT_PROMPT_MAX_BYTES);
722 Self {
723 enabled: true,
724 bundle,
725 max_prompt_items,
726 max_prompt_bytes,
727 include_exclusion_summary: true,
728 include_validation_commands: true,
729 }
730 }
731
732 pub const fn disabled(bundle: SemanticContextBundle) -> Self {
733 Self {
734 enabled: false,
735 bundle,
736 max_prompt_items: DEFAULT_SEMANTIC_CONTEXT_PROMPT_MAX_ITEMS,
737 max_prompt_bytes: DEFAULT_SEMANTIC_CONTEXT_PROMPT_MAX_BYTES,
738 include_exclusion_summary: true,
739 include_validation_commands: true,
740 }
741 }
742
743 #[must_use]
744 pub const fn with_prompt_budget(mut self, max_items: usize, max_bytes: u64) -> Self {
745 self.max_prompt_items = max_items;
746 self.max_prompt_bytes = max_bytes;
747 self
748 }
749}
750
751#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
752#[serde(rename_all = "snake_case")]
753pub enum SemanticContextPromptShape {
754 CustomUserMessage,
755 SystemPromptAppend,
756}
757
758#[derive(Debug, Clone)]
759struct PreparedSemanticContextPrompt {
760 prompt: String,
761 revision: String,
762 shape: SemanticContextPromptShape,
763 details: Value,
764}
765
766#[derive(Debug, Clone, Copy)]
767struct SemanticContextPromptBudget {
768 max_items: usize,
769 max_bytes: u64,
770}
771
772#[derive(Debug, Default, Clone, Copy)]
773struct SemanticContextPromptStats {
774 selected_items_included: usize,
775 selected_items_omitted: usize,
776 validation_commands_included: usize,
777 validation_commands_omitted: usize,
778 exclusions_included: usize,
779 exclusions_omitted: usize,
780 truncated: bool,
781}
782
783pub type MessageFetcher = Arc<dyn Fn() -> BoxFuture<'static, Vec<Message>> + Send + Sync + 'static>;
785
786type AgentEventHandler = Arc<dyn Fn(AgentEvent) + Send + Sync + 'static>;
787
788#[derive(Debug, Clone, Copy, PartialEq, Eq)]
789pub enum QueueMode {
790 All,
791 OneAtATime,
792}
793
794impl QueueMode {
795 pub const fn as_str(self) -> &'static str {
796 match self {
797 Self::All => "all",
798 Self::OneAtATime => "one-at-a-time",
799 }
800 }
801}
802
803#[derive(Debug, Clone, Copy, PartialEq, Eq)]
804pub enum InputSource {
805 Interactive,
806 Rpc,
807 Extension,
808}
809
810impl InputSource {
811 pub const fn as_str(self) -> &'static str {
812 match self {
813 Self::Interactive => "interactive",
814 Self::Rpc => "rpc",
815 Self::Extension => "extension",
816 }
817 }
818}
819
820#[derive(Debug, Clone, Copy)]
821enum QueueKind {
822 Steering,
823 FollowUp,
824}
825
826#[derive(Debug, Clone)]
827struct QueuedMessage {
828 seq: u64,
829 enqueued_at: i64,
830 message: Message,
831}
832
833#[derive(Debug)]
834struct MessageQueue {
835 steering: VecDeque<QueuedMessage>,
836 follow_up: VecDeque<QueuedMessage>,
837 steering_mode: QueueMode,
838 follow_up_mode: QueueMode,
839 next_seq: u64,
840}
841
842impl MessageQueue {
843 const fn new(steering_mode: QueueMode, follow_up_mode: QueueMode) -> Self {
844 Self {
845 steering: VecDeque::new(),
846 follow_up: VecDeque::new(),
847 steering_mode,
848 follow_up_mode,
849 next_seq: 0,
850 }
851 }
852
853 const fn set_modes(&mut self, steering_mode: QueueMode, follow_up_mode: QueueMode) {
854 self.steering_mode = steering_mode;
855 self.follow_up_mode = follow_up_mode;
856 }
857
858 fn pending_count(&self) -> usize {
859 self.steering.len() + self.follow_up.len()
860 }
861
862 fn push(&mut self, kind: QueueKind, message: Message) -> u64 {
863 let seq = self.next_seq;
864 self.next_seq = self.next_seq.saturating_add(1);
865 let entry = QueuedMessage {
866 seq,
867 enqueued_at: Utc::now().timestamp_millis(),
868 message,
869 };
870 match kind {
871 QueueKind::Steering => {
872 if self.steering.len() >= MAX_STEERING_QUEUE_SIZE {
873 tracing::warn!(
874 "Steering queue full ({} messages), dropping oldest message",
875 MAX_STEERING_QUEUE_SIZE
876 );
877 self.steering.pop_front();
878 }
879 self.steering.push_back(entry);
880 }
881 QueueKind::FollowUp => {
882 if self.follow_up.len() >= MAX_FOLLOW_UP_QUEUE_SIZE {
883 tracing::warn!(
884 "Follow-up queue full ({} messages), dropping oldest message",
885 MAX_FOLLOW_UP_QUEUE_SIZE
886 );
887 self.follow_up.pop_front();
888 }
889 self.follow_up.push_back(entry);
890 }
891 }
892 seq
893 }
894
895 fn push_steering(&mut self, message: Message) -> u64 {
896 self.push(QueueKind::Steering, message)
897 }
898
899 fn push_follow_up(&mut self, message: Message) -> u64 {
900 self.push(QueueKind::FollowUp, message)
901 }
902
903 fn pop_steering(&mut self) -> Vec<Message> {
904 self.pop_kind(QueueKind::Steering)
905 }
906
907 fn pop_follow_up(&mut self) -> Vec<Message> {
908 self.pop_kind(QueueKind::FollowUp)
909 }
910
911 fn pop_kind(&mut self, kind: QueueKind) -> Vec<Message> {
912 let (queue, mode) = match kind {
913 QueueKind::Steering => (&mut self.steering, self.steering_mode),
914 QueueKind::FollowUp => (&mut self.follow_up, self.follow_up_mode),
915 };
916
917 match mode {
918 QueueMode::All => queue.drain(..).map(|entry| entry.message).collect(),
919 QueueMode::OneAtATime => queue
920 .pop_front()
921 .into_iter()
922 .map(|entry| entry.message)
923 .collect(),
924 }
925 }
926}
927
928#[derive(Debug, Clone, Serialize)]
934#[serde(tag = "type", rename_all = "snake_case")]
935pub enum AgentEvent {
936 AgentStart {
938 #[serde(rename = "sessionId")]
939 session_id: Arc<str>,
940 },
941 AgentEnd {
943 #[serde(rename = "sessionId")]
944 session_id: Arc<str>,
945 messages: Vec<Message>,
946 #[serde(skip_serializing_if = "Option::is_none")]
947 error: Option<String>,
948 },
949 TurnStart {
951 #[serde(rename = "sessionId")]
952 session_id: Arc<str>,
953 #[serde(rename = "turnIndex")]
954 turn_index: usize,
955 timestamp: i64,
956 },
957 TurnEnd {
959 #[serde(rename = "sessionId")]
960 session_id: Arc<str>,
961 #[serde(rename = "turnIndex")]
962 turn_index: usize,
963 message: Message,
964 #[serde(rename = "toolResults")]
965 tool_results: Vec<Message>,
966 #[serde(rename = "latencyBreakdown", skip_serializing_if = "Option::is_none")]
967 latency_breakdown: Option<Box<TurnLatencyBreakdown>>,
968 },
969 MessageStart { message: Message },
971 MessageUpdate {
973 message: Message,
974 #[serde(rename = "assistantMessageEvent")]
975 assistant_message_event: AssistantMessageEvent,
976 },
977 MessageEnd { message: Message },
979 ToolExecutionStart {
981 #[serde(rename = "toolCallId")]
982 tool_call_id: String,
983 #[serde(rename = "toolName")]
984 tool_name: String,
985 args: serde_json::Value,
986 },
987 ToolExecutionUpdate {
989 #[serde(rename = "toolCallId")]
990 tool_call_id: String,
991 #[serde(rename = "toolName")]
992 tool_name: String,
993 args: serde_json::Value,
994 #[serde(rename = "partialResult")]
995 partial_result: ToolOutput,
996 },
997 ToolExecutionEnd {
999 #[serde(rename = "toolCallId")]
1000 tool_call_id: String,
1001 #[serde(rename = "toolName")]
1002 tool_name: String,
1003 result: ToolOutput,
1004 #[serde(rename = "isError")]
1005 is_error: bool,
1006 },
1007 AutoCompactionStart { reason: String },
1009 AutoCompactionEnd {
1011 #[serde(skip_serializing_if = "Option::is_none")]
1012 result: Option<serde_json::Value>,
1013 aborted: bool,
1014 #[serde(rename = "willRetry")]
1015 will_retry: bool,
1016 #[serde(rename = "errorMessage", skip_serializing_if = "Option::is_none")]
1017 error_message: Option<String>,
1018 },
1019 AutoRetryStart {
1021 attempt: u32,
1022 #[serde(rename = "maxAttempts")]
1023 max_attempts: u32,
1024 #[serde(rename = "delayMs")]
1025 delay_ms: u64,
1026 #[serde(rename = "errorMessage")]
1027 error_message: String,
1028 },
1029 AutoRetryEnd {
1031 success: bool,
1032 attempt: u32,
1033 #[serde(rename = "finalError", skip_serializing_if = "Option::is_none")]
1034 final_error: Option<String>,
1035 },
1036 ExtensionError {
1038 #[serde(rename = "extensionId", skip_serializing_if = "Option::is_none")]
1039 extension_id: Option<String>,
1040 event: String,
1041 error: String,
1042 },
1043}
1044
1045#[derive(Debug, Clone)]
1051pub struct AbortHandle {
1052 inner: Arc<AbortSignalInner>,
1053}
1054
1055#[derive(Debug, Clone)]
1057pub struct AbortSignal {
1058 inner: Arc<AbortSignalInner>,
1059}
1060
1061#[derive(Debug)]
1062struct AbortSignalInner {
1063 aborted: AtomicBool,
1064 notify: Notify,
1065}
1066
1067impl AbortHandle {
1068 #[must_use]
1070 pub fn new() -> (Self, AbortSignal) {
1071 let inner = Arc::new(AbortSignalInner {
1072 aborted: AtomicBool::new(false),
1073 notify: Notify::new(),
1074 });
1075 (
1076 Self {
1077 inner: Arc::clone(&inner),
1078 },
1079 AbortSignal { inner },
1080 )
1081 }
1082
1083 pub fn abort(&self) {
1085 if !self.inner.aborted.swap(true, Ordering::SeqCst) {
1086 self.inner.notify.notify_waiters();
1087 }
1088 }
1089}
1090
1091impl AbortSignal {
1092 #[must_use]
1094 pub fn is_aborted(&self) -> bool {
1095 self.inner.aborted.load(Ordering::SeqCst)
1096 }
1097
1098 pub async fn wait(&self) {
1099 if self.is_aborted() {
1100 return;
1101 }
1102
1103 loop {
1104 self.inner.notify.notified().await;
1105 if self.is_aborted() {
1106 return;
1107 }
1108 }
1109 }
1110}
1111
1112pub struct Agent {
1114 provider: Arc<dyn Provider>,
1116
1117 tools: ToolRegistry,
1119
1120 config: AgentConfig,
1122
1123 extensions: Option<ExtensionManager>,
1125
1126 messages: Vec<Message>,
1128
1129 steering_fetchers: Vec<MessageFetcher>,
1131
1132 follow_up_fetchers: Vec<MessageFetcher>,
1134
1135 message_queue: MessageQueue,
1137
1138 cached_tool_defs: Option<Vec<ToolDef>>,
1140}
1141
1142impl Agent {
1143 pub fn new(provider: Arc<dyn Provider>, tools: ToolRegistry, config: AgentConfig) -> Self {
1145 Self {
1146 provider,
1147 tools,
1148 config,
1149 extensions: None,
1150 messages: Vec::new(),
1151 steering_fetchers: Vec::new(),
1152 follow_up_fetchers: Vec::new(),
1153 message_queue: MessageQueue::new(QueueMode::OneAtATime, QueueMode::OneAtATime),
1154 cached_tool_defs: None,
1155 }
1156 }
1157
1158 #[must_use]
1160 pub fn messages(&self) -> &[Message] {
1161 &self.messages
1162 }
1163
1164 pub fn clear_messages(&mut self) {
1166 self.messages.clear();
1167 }
1168
1169 pub fn add_message(&mut self, message: Message) {
1171 if self.messages.len() >= MAX_AGENT_MESSAGES {
1172 tracing::warn!(
1173 "Agent message history full ({} messages), dropping oldest message",
1174 MAX_AGENT_MESSAGES
1175 );
1176 self.messages.remove(0);
1177 }
1178 self.messages.push(message);
1179 }
1180
1181 pub fn replace_messages(&mut self, messages: Vec<Message>) {
1183 self.messages = messages;
1184 }
1185
1186 pub fn set_provider(&mut self, provider: Arc<dyn Provider>) {
1188 self.provider = provider;
1189 }
1190
1191 pub fn register_message_fetchers(
1196 &mut self,
1197 steering: Option<MessageFetcher>,
1198 follow_up: Option<MessageFetcher>,
1199 ) {
1200 if let Some(fetcher) = steering {
1201 self.steering_fetchers.push(fetcher);
1202 }
1203 if let Some(fetcher) = follow_up {
1204 self.follow_up_fetchers.push(fetcher);
1205 }
1206 }
1207
1208 pub fn extend_tools<I>(&mut self, tools: I)
1210 where
1211 I: IntoIterator<Item = Box<dyn Tool>>,
1212 {
1213 self.tools.extend(tools);
1214 self.cached_tool_defs = None; }
1216
1217 pub fn queue_steering(&mut self, message: Message) -> u64 {
1219 self.message_queue.push_steering(message)
1220 }
1221
1222 pub fn queue_follow_up(&mut self, message: Message) -> u64 {
1224 self.message_queue.push_follow_up(message)
1225 }
1226
1227 pub const fn set_queue_modes(&mut self, steering: QueueMode, follow_up: QueueMode) {
1229 self.message_queue.set_modes(steering, follow_up);
1230 }
1231
1232 pub const fn queue_modes(&self) -> (QueueMode, QueueMode) {
1233 (
1234 self.message_queue.steering_mode,
1235 self.message_queue.follow_up_mode,
1236 )
1237 }
1238
1239 #[must_use]
1241 pub fn queued_message_count(&self) -> usize {
1242 self.message_queue.pending_count()
1243 }
1244
1245 pub fn provider(&self) -> Arc<dyn Provider> {
1246 Arc::clone(&self.provider)
1247 }
1248
1249 pub const fn stream_options(&self) -> &StreamOptions {
1250 &self.config.stream_options
1251 }
1252
1253 pub const fn stream_options_mut(&mut self) -> &mut StreamOptions {
1254 &mut self.config.stream_options
1255 }
1256
1257 pub fn system_prompt(&self) -> Option<&str> {
1258 self.config.system_prompt.as_deref()
1259 }
1260
1261 pub fn set_system_prompt(&mut self, system_prompt: Option<String>) {
1262 self.config.system_prompt = system_prompt;
1263 }
1264
1265 fn build_context(&mut self) -> Context<'_> {
1267 let messages: Cow<'_, [Message]> = if self.config.block_images {
1268 let mut msgs = self.messages.clone();
1269 msgs.retain(|m| match m {
1271 Message::Custom(c) => c.display,
1272 _ => true,
1273 });
1274 let stats = filter_images_for_provider(&mut msgs);
1275 if stats.removed_images > 0 {
1276 tracing::debug!(
1277 filtered_images = stats.removed_images,
1278 affected_messages = stats.affected_messages,
1279 "Filtered image content from outbound provider context (images.block_images=true)"
1280 );
1281 }
1282 Cow::Owned(msgs)
1283 } else {
1284 let has_hidden = self.messages.iter().any(|m| match m {
1286 Message::Custom(c) => !c.display,
1287 _ => false,
1288 });
1289
1290 if has_hidden {
1291 let mut msgs = self.messages.clone();
1292 msgs.retain(|m| match m {
1293 Message::Custom(c) => c.display,
1294 _ => true,
1295 });
1296 Cow::Owned(msgs)
1297 } else {
1298 Cow::Borrowed(self.messages.as_slice())
1299 }
1300 };
1301
1302 if self.cached_tool_defs.is_none() {
1304 let defs: Vec<ToolDef> = self
1305 .tools
1306 .tools()
1307 .iter()
1308 .map(|t| ToolDef {
1309 name: t.name().to_string(),
1310 description: t.description().to_string(),
1311 parameters: t.parameters(),
1312 })
1313 .collect();
1314 self.cached_tool_defs = Some(defs);
1315 }
1316 let tools = Cow::Borrowed(self.cached_tool_defs.as_deref().unwrap());
1317
1318 Context {
1319 system_prompt: self.config.system_prompt.as_deref().map(Cow::Borrowed),
1320 messages,
1321 tools,
1322 }
1323 }
1324
1325 pub async fn run(
1329 &mut self,
1330 user_input: impl Into<String>,
1331 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
1332 ) -> Result<AssistantMessage> {
1333 self.run_with_abort(user_input, None, on_event).await
1334 }
1335
1336 pub async fn run_with_abort(
1338 &mut self,
1339 user_input: impl Into<String>,
1340 abort: Option<AbortSignal>,
1341 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
1342 ) -> Result<AssistantMessage> {
1343 let user_message = Message::User(UserMessage {
1345 content: UserContent::Text(user_input.into()),
1346 timestamp: Utc::now().timestamp_millis(),
1347 });
1348
1349 self.run_loop(vec![user_message], Arc::new(on_event), abort)
1351 .await
1352 }
1353
1354 pub async fn run_with_content(
1356 &mut self,
1357 content: Vec<ContentBlock>,
1358 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
1359 ) -> Result<AssistantMessage> {
1360 self.run_with_content_with_abort(content, None, on_event)
1361 .await
1362 }
1363
1364 pub async fn run_with_content_with_abort(
1366 &mut self,
1367 content: Vec<ContentBlock>,
1368 abort: Option<AbortSignal>,
1369 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
1370 ) -> Result<AssistantMessage> {
1371 let user_message = Message::User(UserMessage {
1373 content: UserContent::Blocks(content),
1374 timestamp: Utc::now().timestamp_millis(),
1375 });
1376
1377 self.run_loop(vec![user_message], Arc::new(on_event), abort)
1379 .await
1380 }
1381
1382 pub async fn run_with_message_with_abort(
1384 &mut self,
1385 message: Message,
1386 abort: Option<AbortSignal>,
1387 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
1388 ) -> Result<AssistantMessage> {
1389 self.run_loop(vec![message], Arc::new(on_event), abort)
1390 .await
1391 }
1392
1393 pub async fn run_with_messages_with_abort(
1395 &mut self,
1396 messages: Vec<Message>,
1397 abort: Option<AbortSignal>,
1398 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
1399 ) -> Result<AssistantMessage> {
1400 self.run_loop(messages, Arc::new(on_event), abort).await
1401 }
1402
1403 pub async fn run_continue_with_abort(
1405 &mut self,
1406 abort: Option<AbortSignal>,
1407 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
1408 ) -> Result<AssistantMessage> {
1409 self.run_loop(Vec::new(), Arc::new(on_event), abort).await
1410 }
1411
1412 fn build_abort_message(&self, partial: Option<&AssistantMessage>) -> AssistantMessage {
1413 let mut message = partial.cloned().unwrap_or_else(|| AssistantMessage {
1414 content: Vec::new(),
1415 api: self.provider.api().to_string(),
1416 provider: self.provider.name().to_string(),
1417 model: self.provider.model_id().to_string(),
1418 usage: Usage::default(),
1419 stop_reason: StopReason::Aborted,
1420 error_message: Some("Aborted".to_string()),
1421 timestamp: Utc::now().timestamp_millis(),
1422 });
1423 message.stop_reason = StopReason::Aborted;
1424 message.error_message = Some("Aborted".to_string());
1425 message.timestamp = Utc::now().timestamp_millis();
1426 message
1427 }
1428
1429 fn build_error_message(
1430 &self,
1431 partial: Option<&AssistantMessage>,
1432 error_message: impl Into<String>,
1433 ) -> AssistantMessage {
1434 let error_message = error_message.into();
1435 let mut message = partial.cloned().unwrap_or_else(|| AssistantMessage {
1436 content: Vec::new(),
1437 api: self.provider.api().to_string(),
1438 provider: self.provider.name().to_string(),
1439 model: self.provider.model_id().to_string(),
1440 usage: Usage::default(),
1441 stop_reason: StopReason::Error,
1442 error_message: Some(error_message.clone()),
1443 timestamp: Utc::now().timestamp_millis(),
1444 });
1445 message.stop_reason = StopReason::Error;
1446 message.error_message = Some(error_message);
1447 message.timestamp = Utc::now().timestamp_millis();
1448 message
1449 }
1450
1451 #[allow(clippy::too_many_lines)]
1453 async fn run_loop(
1454 &mut self,
1455 prompts: Vec<Message>,
1456 on_event: AgentEventHandler,
1457 abort: Option<AbortSignal>,
1458 ) -> Result<AssistantMessage> {
1459 let loop_cx = crate::agent_cx::AgentCx::for_current_or_request();
1460 let session_id: Arc<str> = self
1461 .config
1462 .stream_options
1463 .session_id
1464 .as_deref()
1465 .unwrap_or("")
1466 .into();
1467 let mut iterations = 0usize;
1468 let mut warned_at_handoff_threshold = false;
1469 let mut turn_index: usize = 0;
1470 let mut new_messages: Vec<Message> = Vec::with_capacity(prompts.len() + 8);
1471 let mut last_assistant: Option<Arc<AssistantMessage>> = None;
1472
1473 let agent_start_event = AgentEvent::AgentStart {
1474 session_id: session_id.clone(),
1475 };
1476 self.dispatch_extension_lifecycle_event(&agent_start_event)
1477 .await;
1478 on_event(agent_start_event);
1479
1480 for prompt in prompts {
1481 self.messages.push(prompt.clone());
1482 on_event(AgentEvent::MessageStart {
1483 message: prompt.clone(),
1484 });
1485 on_event(AgentEvent::MessageEnd {
1486 message: prompt.clone(),
1487 });
1488 new_messages.push(prompt);
1489 }
1490
1491 let mut pending_messages = self.drain_steering_messages().await;
1493
1494 loop {
1495 let mut has_more_tool_calls = true;
1496 let mut steering_after_tools: Option<Vec<Message>> = None;
1497
1498 while has_more_tool_calls || !pending_messages.is_empty() {
1499 let current_turn_index = turn_index;
1500 let turn_latency = Arc::new(StdMutex::new(TurnLatencyAccumulator::started()));
1501 let turn_start_event = AgentEvent::TurnStart {
1502 session_id: session_id.clone(),
1503 turn_index: current_turn_index,
1504 timestamp: Utc::now().timestamp_millis(),
1505 };
1506 self.dispatch_extension_lifecycle_event(&turn_start_event)
1507 .await;
1508 on_event(turn_start_event);
1509
1510 for message in std::mem::take(&mut pending_messages) {
1511 self.messages.push(message.clone());
1512 on_event(AgentEvent::MessageStart {
1513 message: message.clone(),
1514 });
1515 on_event(AgentEvent::MessageEnd {
1516 message: message.clone(),
1517 });
1518 new_messages.push(message);
1519 }
1520
1521 if abort.as_ref().is_some_and(AbortSignal::is_aborted) {
1522 let abort_message = self.build_abort_message(None);
1523 let message = Message::assistant(abort_message.clone());
1524
1525 self.messages.push(message.clone());
1526 new_messages.push(message.clone());
1527 on_event(AgentEvent::MessageStart {
1528 message: message.clone(),
1529 });
1530 on_event(AgentEvent::MessageEnd {
1531 message: message.clone(),
1532 });
1533
1534 let turn_end_event = AgentEvent::TurnEnd {
1535 session_id: session_id.clone(),
1536 turn_index: current_turn_index,
1537 message,
1538 tool_results: Vec::new(),
1539 latency_breakdown: snapshot_turn_latency(&turn_latency),
1540 };
1541 self.dispatch_extension_lifecycle_event(&turn_end_event)
1542 .await;
1543 on_event(turn_end_event);
1544 let agent_end_event = AgentEvent::AgentEnd {
1545 session_id: session_id.clone(),
1546 messages: std::mem::take(&mut new_messages),
1547 error: Some(
1548 abort_message
1549 .error_message
1550 .clone()
1551 .unwrap_or_else(|| "Aborted".to_string()),
1552 ),
1553 };
1554 self.dispatch_extension_lifecycle_event(&agent_end_event)
1555 .await;
1556 on_event(agent_end_event);
1557 return Ok(abort_message);
1558 }
1559
1560 let provider_streaming_started_at = Instant::now();
1561 let assistant_result = self
1562 .stream_assistant_response(Arc::clone(&on_event), abort.clone(), &loop_cx)
1563 .await;
1564 record_provider_streaming_latency(
1565 &turn_latency,
1566 provider_streaming_started_at.elapsed(),
1567 );
1568
1569 let assistant_message = match assistant_result {
1570 Ok(msg) => msg,
1571 Err(err) => {
1572 let err_string = err.to_string();
1573 let steering_to_add = self.drain_steering_messages().await;
1574 for message in steering_to_add {
1575 self.messages.push(message.clone());
1576 on_event(AgentEvent::MessageStart {
1577 message: message.clone(),
1578 });
1579 on_event(AgentEvent::MessageEnd {
1580 message: message.clone(),
1581 });
1582 new_messages.push(message);
1583 }
1584
1585 let error_message = self.build_error_message(None, err_string.clone());
1586 let assistant_event_message = Message::assistant(error_message.clone());
1587 self.messages.push(assistant_event_message.clone());
1588 new_messages.push(assistant_event_message.clone());
1589 on_event(AgentEvent::MessageStart {
1590 message: assistant_event_message.clone(),
1591 });
1592 on_event(AgentEvent::MessageEnd {
1593 message: assistant_event_message.clone(),
1594 });
1595
1596 let turn_end_event = AgentEvent::TurnEnd {
1597 session_id: session_id.clone(),
1598 turn_index: current_turn_index,
1599 message: assistant_event_message,
1600 tool_results: Vec::new(),
1601 latency_breakdown: snapshot_turn_latency(&turn_latency),
1602 };
1603 self.dispatch_extension_lifecycle_event(&turn_end_event)
1604 .await;
1605 on_event(turn_end_event);
1606
1607 let agent_end_event = AgentEvent::AgentEnd {
1608 session_id: session_id.clone(),
1609 messages: std::mem::take(&mut new_messages),
1610 error: Some(err_string),
1611 };
1612 self.dispatch_extension_lifecycle_event(&agent_end_event)
1613 .await;
1614 on_event(agent_end_event);
1615 return Err(err);
1616 }
1617 };
1618 let assistant_arc = Arc::new(assistant_message);
1621 last_assistant = Some(Arc::clone(&assistant_arc));
1622
1623 let assistant_event_message = Message::Assistant(Arc::clone(&assistant_arc));
1624 new_messages.push(assistant_event_message.clone());
1625
1626 if matches!(
1627 assistant_arc.stop_reason,
1628 StopReason::Error | StopReason::Aborted
1629 ) {
1630 let steering_to_add = self.drain_steering_messages().await;
1631 for message in steering_to_add {
1632 self.messages.push(message.clone());
1633 on_event(AgentEvent::MessageStart {
1634 message: message.clone(),
1635 });
1636 on_event(AgentEvent::MessageEnd {
1637 message: message.clone(),
1638 });
1639 new_messages.push(message);
1640 }
1641
1642 let turn_end_event = AgentEvent::TurnEnd {
1643 session_id: session_id.clone(),
1644 turn_index: current_turn_index,
1645 message: assistant_event_message.clone(),
1646 tool_results: Vec::new(),
1647 latency_breakdown: snapshot_turn_latency(&turn_latency),
1648 };
1649 self.dispatch_extension_lifecycle_event(&turn_end_event)
1650 .await;
1651 on_event(turn_end_event);
1652 let agent_end_event = AgentEvent::AgentEnd {
1653 session_id: session_id.clone(),
1654 messages: std::mem::take(&mut new_messages),
1655 error: assistant_arc.error_message.clone(),
1656 };
1657 self.dispatch_extension_lifecycle_event(&agent_end_event)
1658 .await;
1659 on_event(agent_end_event);
1660 return Ok(Arc::unwrap_or_clone(assistant_arc));
1661 }
1662
1663 let tool_calls = extract_tool_calls(&assistant_arc.content);
1664 has_more_tool_calls = !tool_calls.is_empty();
1665
1666 let mut tool_results: Vec<Arc<ToolResultMessage>> = Vec::new();
1667 if has_more_tool_calls {
1668 iterations += 1;
1669 if !warned_at_handoff_threshold
1677 && should_warn_at_iteration_threshold(
1678 iterations,
1679 self.config.max_tool_iterations,
1680 )
1681 {
1682 warned_at_handoff_threshold = true;
1683 let warning = Message::User(UserMessage {
1684 content: UserContent::Text(iteration_handoff_steering_text(
1685 iterations,
1686 self.config.max_tool_iterations,
1687 )),
1688 timestamp: Utc::now().timestamp_millis(),
1689 });
1690 self.message_queue.push_steering(warning);
1691 tracing::warn!(
1692 iterations,
1693 max = self.config.max_tool_iterations,
1694 "tool-iteration budget at >=80%; injected handoff steering message"
1695 );
1696 }
1697 if iterations > self.config.max_tool_iterations {
1698 let error_message = format!(
1699 "Maximum tool iterations ({}) exceeded",
1700 self.config.max_tool_iterations
1701 );
1702 let mut stop_message = (*assistant_arc).clone();
1703 stop_message.stop_reason = StopReason::Error;
1704 stop_message.error_message = Some(error_message.clone());
1705
1706 stop_message
1708 .content
1709 .retain(|b| !matches!(b, crate::model::ContentBlock::ToolCall(_)));
1710
1711 let stop_arc = Arc::new(stop_message.clone());
1712 let stop_event_message = Message::Assistant(Arc::clone(&stop_arc));
1713
1714 if let Some(last @ Message::Assistant(_)) = self
1717 .messages
1718 .iter_mut()
1719 .rev()
1720 .find(|m| matches!(m, Message::Assistant(_)))
1721 {
1722 *last = stop_event_message.clone();
1723 }
1724 if let Some(last @ Message::Assistant(_)) = new_messages.last_mut() {
1725 *last = stop_event_message.clone();
1726 }
1727
1728 let steering_to_add = self.drain_steering_messages().await;
1729 for message in steering_to_add {
1730 self.messages.push(message.clone());
1731 on_event(AgentEvent::MessageStart {
1732 message: message.clone(),
1733 });
1734 on_event(AgentEvent::MessageEnd {
1735 message: message.clone(),
1736 });
1737 new_messages.push(message);
1738 }
1739
1740 let turn_end_event = AgentEvent::TurnEnd {
1741 session_id: session_id.clone(),
1742 turn_index: current_turn_index,
1743 message: stop_event_message,
1744 tool_results: Vec::new(),
1745 latency_breakdown: snapshot_turn_latency(&turn_latency),
1746 };
1747 self.dispatch_extension_lifecycle_event(&turn_end_event)
1748 .await;
1749 on_event(turn_end_event);
1750
1751 let agent_end_event = AgentEvent::AgentEnd {
1752 session_id: session_id.clone(),
1753 messages: std::mem::take(&mut new_messages),
1754 error: Some(error_message),
1755 };
1756 self.dispatch_extension_lifecycle_event(&agent_end_event)
1757 .await;
1758 on_event(agent_end_event);
1759
1760 return Ok(stop_message);
1761 }
1762
1763 let outcome = match self
1764 .execute_tool_calls(
1765 &tool_calls,
1766 Arc::clone(&on_event),
1767 &mut new_messages,
1768 abort.clone(),
1769 Arc::clone(&turn_latency),
1770 )
1771 .await
1772 {
1773 Ok(outcome) => outcome,
1774 Err(err) => {
1775 let steering_to_add = self.drain_steering_messages().await;
1776 for message in steering_to_add {
1777 self.messages.push(message.clone());
1778 on_event(AgentEvent::MessageStart {
1779 message: message.clone(),
1780 });
1781 on_event(AgentEvent::MessageEnd {
1782 message: message.clone(),
1783 });
1784 new_messages.push(message);
1785 }
1786
1787 let turn_end_event = AgentEvent::TurnEnd {
1788 session_id: session_id.clone(),
1789 turn_index: current_turn_index,
1790 message: assistant_event_message.clone(),
1791 tool_results: Vec::new(),
1792 latency_breakdown: snapshot_turn_latency(&turn_latency),
1793 };
1794 self.dispatch_extension_lifecycle_event(&turn_end_event)
1795 .await;
1796 on_event(turn_end_event);
1797
1798 let agent_end_event = AgentEvent::AgentEnd {
1799 session_id: session_id.clone(),
1800 messages: std::mem::take(&mut new_messages),
1801 error: Some(err.to_string()),
1802 };
1803 self.dispatch_extension_lifecycle_event(&agent_end_event)
1804 .await;
1805 on_event(agent_end_event);
1806 return Err(err);
1807 }
1808 };
1809 tool_results = outcome.tool_results;
1810 steering_after_tools = outcome.steering_messages;
1811 }
1812
1813 let tool_messages = tool_results
1814 .iter()
1815 .map(|r| Message::ToolResult(Arc::clone(r)))
1816 .collect::<Vec<_>>();
1817
1818 let turn_end_event = AgentEvent::TurnEnd {
1819 session_id: session_id.clone(),
1820 turn_index: current_turn_index,
1821 message: assistant_event_message.clone(),
1822 tool_results: tool_messages,
1823 latency_breakdown: snapshot_turn_latency(&turn_latency),
1824 };
1825 self.dispatch_extension_lifecycle_event(&turn_end_event)
1826 .await;
1827 on_event(turn_end_event);
1828
1829 turn_index = turn_index.saturating_add(1);
1830
1831 if let Some(steering) = steering_after_tools.take() {
1832 pending_messages = steering;
1833 } else {
1834 pending_messages = self.drain_steering_messages().await;
1836 }
1837 }
1838
1839 let follow_up = self.drain_follow_up_messages().await;
1841 if follow_up.is_empty() {
1842 break;
1843 }
1844 pending_messages = follow_up;
1845 }
1846
1847 let Some(final_arc) = last_assistant else {
1848 return Err(Error::api("Agent completed without assistant message"));
1849 };
1850
1851 let agent_end_event = AgentEvent::AgentEnd {
1852 session_id: session_id.clone(),
1853 messages: new_messages,
1854 error: None,
1855 };
1856 self.dispatch_extension_lifecycle_event(&agent_end_event)
1857 .await;
1858 on_event(agent_end_event);
1859 Ok(Arc::unwrap_or_clone(final_arc))
1860 }
1861
1862 async fn fetch_messages(&self, fetcher: Option<&MessageFetcher>) -> Vec<Message> {
1863 if let Some(fetcher) = fetcher {
1864 (fetcher)().await
1865 } else {
1866 Vec::new()
1867 }
1868 }
1869
1870 async fn dispatch_extension_lifecycle_event(&self, event: &AgentEvent) {
1871 let Some(extensions) = &self.extensions else {
1872 return;
1873 };
1874
1875 let name = match event {
1876 AgentEvent::AgentStart { .. } => ExtensionEventName::AgentStart,
1877 AgentEvent::AgentEnd { .. } => ExtensionEventName::AgentEnd,
1878 AgentEvent::TurnStart { .. } => ExtensionEventName::TurnStart,
1879 AgentEvent::TurnEnd { .. } => ExtensionEventName::TurnEnd,
1880 _ => return,
1881 };
1882
1883 let payload = match serde_json::to_value(event) {
1884 Ok(payload) => payload,
1885 Err(err) => {
1886 tracing::warn!("failed to serialize agent lifecycle event (fail-open): {err}");
1887 return;
1888 }
1889 };
1890
1891 if let Err(err) = extensions.dispatch_event(name, Some(payload)).await {
1892 tracing::warn!("agent lifecycle extension hook failed (fail-open): {err}");
1893 }
1894 }
1895
1896 async fn dispatch_context_event(&self, messages: &[Message]) -> Option<Vec<Message>> {
1897 let Some(extensions) = &self.extensions else {
1898 return None;
1899 };
1900
1901 let payload = json!({ "messages": messages });
1902 let response = extensions
1903 .dispatch_event_with_response(
1904 ExtensionEventName::Context,
1905 Some(payload),
1906 EXTENSION_EVENT_TIMEOUT_MS,
1907 )
1908 .await
1909 .ok()?;
1910
1911 let value = response?;
1912
1913 if value.is_null() {
1914 return None;
1915 }
1916
1917 let messages_value = if let Some(obj) = value.as_object() {
1918 obj.get("messages").cloned()?
1919 } else if value.is_array() {
1920 value
1921 } else {
1922 return None;
1923 };
1924
1925 if messages_value.is_null() {
1926 return Some(Vec::new());
1927 }
1928
1929 match serde_json::from_value(messages_value) {
1930 Ok(messages) => Some(messages),
1931 Err(err) => {
1932 tracing::warn!("context extension hook returned invalid messages: {err}");
1933 None
1934 }
1935 }
1936 }
1937
1938 async fn drain_steering_messages(&mut self) -> Vec<Message> {
1939 for fetcher in &self.steering_fetchers {
1940 let fetched = self.fetch_messages(Some(fetcher)).await;
1941 for message in fetched {
1942 self.message_queue.push_steering(message);
1943 }
1944 }
1945 self.message_queue.pop_steering()
1946 }
1947
1948 async fn drain_follow_up_messages(&mut self) -> Vec<Message> {
1949 for fetcher in &self.follow_up_fetchers {
1950 let fetched = self.fetch_messages(Some(fetcher)).await;
1951 for message in fetched {
1952 self.message_queue.push_follow_up(message);
1953 }
1954 }
1955 self.message_queue.pop_follow_up()
1956 }
1957
1958 #[allow(clippy::too_many_lines)]
1960 async fn stream_assistant_response(
1961 &mut self,
1962 on_event: AgentEventHandler,
1963 abort: Option<AbortSignal>,
1964 checkpoint_cx: &crate::agent_cx::AgentCx,
1965 ) -> Result<AssistantMessage> {
1966 let provider = Arc::clone(&self.provider);
1968 let stream_options = self.config.stream_options.clone();
1969 let (system_prompt, tools, base_messages) = {
1970 let context = self.build_context();
1971 (
1972 context.system_prompt.as_deref().map(str::to_string),
1973 context.tools.to_vec(),
1974 context.messages.to_vec(),
1975 )
1976 };
1977 let messages = self
1978 .dispatch_context_event(&base_messages)
1979 .await
1980 .unwrap_or(base_messages);
1981 let context = Context::owned(system_prompt, messages, tools);
1982 let mut stream = provider.stream(&context, &stream_options).await?;
1983
1984 let mut added_partial = false;
1985 let mut sent_start = false;
1988 let mut tool_call_raw_args: std::collections::HashMap<usize, String> =
1995 std::collections::HashMap::new();
1996
1997 'stream: loop {
1998 if checkpoint_cx.checkpoint().is_err() {
1999 let last_partial = if added_partial {
2000 match self
2001 .messages
2002 .iter()
2003 .rev()
2004 .find(|m| matches!(m, Message::Assistant(_)))
2005 {
2006 Some(Message::Assistant(a)) => Some(a.as_ref()),
2007 _ => None,
2008 }
2009 } else {
2010 None
2011 };
2012 let abort_arc = Arc::new(self.build_abort_message(last_partial));
2013 if !sent_start {
2014 on_event(AgentEvent::MessageStart {
2015 message: Message::Assistant(Arc::clone(&abort_arc)),
2016 });
2017 self.messages
2018 .push(Message::Assistant(Arc::clone(&abort_arc)));
2019 added_partial = true;
2020 }
2021 on_event(AgentEvent::MessageUpdate {
2022 message: Message::Assistant(Arc::clone(&abort_arc)),
2023 assistant_message_event: AssistantMessageEvent::Error {
2024 reason: StopReason::Aborted,
2025 error: Arc::clone(&abort_arc),
2026 },
2027 });
2028 return Ok(self.finalize_assistant_message(
2029 Arc::try_unwrap(abort_arc).unwrap_or_else(|a| (*a).clone()),
2030 &on_event,
2031 added_partial,
2032 ));
2033 }
2034
2035 let event_result = if let Some(signal) = abort.as_ref() {
2036 let abort_fut = signal.wait().fuse();
2037 let event_fut = stream.next().fuse();
2038 futures::pin_mut!(abort_fut, event_fut);
2039
2040 match futures::future::select(abort_fut, event_fut).await {
2041 futures::future::Either::Left(((), _event_fut)) => {
2042 let last_partial = if added_partial {
2043 match self
2044 .messages
2045 .iter()
2046 .rev()
2047 .find(|m| matches!(m, Message::Assistant(_)))
2048 {
2049 Some(Message::Assistant(a)) => Some(a.as_ref()),
2050 _ => None,
2051 }
2052 } else {
2053 None
2054 };
2055 let abort_arc = Arc::new(self.build_abort_message(last_partial));
2056 if !sent_start {
2057 on_event(AgentEvent::MessageStart {
2058 message: Message::Assistant(Arc::clone(&abort_arc)),
2059 });
2060 self.messages
2061 .push(Message::Assistant(Arc::clone(&abort_arc)));
2062 added_partial = true;
2063 }
2067 on_event(AgentEvent::MessageUpdate {
2068 message: Message::Assistant(Arc::clone(&abort_arc)),
2069 assistant_message_event: AssistantMessageEvent::Error {
2070 reason: StopReason::Aborted,
2071 error: Arc::clone(&abort_arc),
2072 },
2073 });
2074 return Ok(self.finalize_assistant_message(
2075 Arc::try_unwrap(abort_arc).unwrap_or_else(|a| (*a).clone()),
2076 &on_event,
2077 added_partial,
2078 ));
2079 }
2080 futures::future::Either::Right((event, _abort_fut)) => event,
2081 }
2082 } else {
2083 let event_fut = stream.next().fuse();
2084 futures::pin_mut!(event_fut);
2085 loop {
2086 let now = checkpoint_cx
2087 .cx()
2088 .timer_driver()
2089 .map_or_else(asupersync::time::wall_now, |timer| timer.now());
2090 let tick_fut =
2091 asupersync::time::sleep(now, std::time::Duration::from_millis(25)).fuse();
2092 futures::pin_mut!(tick_fut);
2093
2094 match futures::future::select(tick_fut, &mut event_fut).await {
2095 futures::future::Either::Left(((), _event_fut)) => {
2096 if checkpoint_cx.checkpoint().is_err() {
2097 continue 'stream;
2098 }
2099 }
2100 futures::future::Either::Right((result, _tick_fut)) => break result,
2101 }
2102 }
2103 };
2104
2105 let Some(event_result) = event_result else {
2106 break;
2107 };
2108 let event = match event_result {
2109 Ok(e) => e,
2110 Err(err) => {
2111 let partial = if added_partial {
2112 match self
2113 .messages
2114 .iter()
2115 .rev()
2116 .find(|m| matches!(m, Message::Assistant(_)))
2117 {
2118 Some(Message::Assistant(a)) => Some(a.as_ref()),
2119 _ => None,
2120 }
2121 } else {
2122 None
2123 };
2124 let msg = self.build_error_message(partial, err.to_string());
2125
2126 return Ok(self.finalize_assistant_message(msg, &on_event, added_partial));
2130 }
2131 };
2132
2133 match event {
2134 StreamEvent::Start { partial } => {
2135 if added_partial {
2136 if let Some(Message::Assistant(msg_arc)) = self
2137 .messages
2138 .iter_mut()
2139 .rev()
2140 .find(|m| matches!(m, Message::Assistant(_)))
2141 {
2142 let msg = Arc::make_mut(msg_arc);
2143 if msg.content.is_empty() {
2144 *msg = partial;
2145 } else {
2146 msg.api = partial.api;
2147 msg.provider = partial.provider;
2148 msg.model = partial.model;
2149 msg.usage = partial.usage;
2150 msg.stop_reason = partial.stop_reason;
2151 msg.error_message = partial.error_message;
2152 msg.timestamp = partial.timestamp;
2153 }
2154 let shared = Arc::clone(msg_arc);
2155 if !sent_start {
2156 on_event(AgentEvent::MessageStart {
2157 message: Message::Assistant(Arc::clone(&shared)),
2158 });
2159 sent_start = true;
2160 }
2161 on_event(AgentEvent::MessageUpdate {
2162 message: Message::Assistant(Arc::clone(&shared)),
2163 assistant_message_event: AssistantMessageEvent::Start {
2164 partial: shared,
2165 },
2166 });
2167 } else {
2168 let shared = Arc::new(partial);
2169 self.update_partial_message(Arc::clone(&shared), &mut added_partial);
2170 on_event(AgentEvent::MessageStart {
2171 message: Message::Assistant(Arc::clone(&shared)),
2172 });
2173 sent_start = true;
2174 on_event(AgentEvent::MessageUpdate {
2175 message: Message::Assistant(Arc::clone(&shared)),
2176 assistant_message_event: AssistantMessageEvent::Start {
2177 partial: shared,
2178 },
2179 });
2180 }
2181 } else {
2182 let shared = Arc::new(partial);
2183 self.update_partial_message(Arc::clone(&shared), &mut added_partial);
2184 on_event(AgentEvent::MessageStart {
2185 message: Message::Assistant(Arc::clone(&shared)),
2186 });
2187 sent_start = true;
2188 on_event(AgentEvent::MessageUpdate {
2189 message: Message::Assistant(Arc::clone(&shared)),
2190 assistant_message_event: AssistantMessageEvent::Start {
2191 partial: shared,
2192 },
2193 });
2194 }
2195 }
2196 StreamEvent::TextStart { content_index, .. } => {
2197 self.seed_partial_message_if_missing(&mut added_partial);
2198 if let Some(Message::Assistant(msg_arc)) = self
2199 .messages
2200 .iter_mut()
2201 .rev()
2202 .find(|m| matches!(m, Message::Assistant(_)))
2203 {
2204 let msg = Arc::make_mut(msg_arc);
2205 if content_index == msg.content.len() {
2206 msg.content.push(ContentBlock::Text(TextContent::new("")));
2207 }
2208 let shared = Arc::clone(msg_arc);
2209 if !sent_start {
2210 on_event(AgentEvent::MessageStart {
2211 message: Message::Assistant(Arc::clone(&shared)),
2212 });
2213 sent_start = true;
2214 }
2215 on_event(AgentEvent::MessageUpdate {
2216 message: Message::Assistant(Arc::clone(&shared)),
2217 assistant_message_event: AssistantMessageEvent::TextStart {
2218 content_index,
2219 partial: shared,
2220 },
2221 });
2222 }
2223 }
2224 StreamEvent::TextDelta {
2225 content_index,
2226 delta,
2227 ..
2228 } => {
2229 self.seed_partial_message_if_missing(&mut added_partial);
2230 if let Some(Message::Assistant(msg_arc)) = self
2231 .messages
2232 .iter_mut()
2233 .rev()
2234 .find(|m| matches!(m, Message::Assistant(_)))
2235 {
2236 {
2237 let msg = Arc::make_mut(msg_arc);
2238 if msg.content.get(content_index).is_none()
2239 && content_index == msg.content.len()
2240 {
2241 msg.content.push(ContentBlock::Text(TextContent::new("")));
2242 }
2243 if let Some(ContentBlock::Text(text)) =
2244 msg.content.get_mut(content_index)
2245 {
2246 text.text.push_str(&delta);
2247 }
2248 }
2249 let shared = Arc::clone(msg_arc);
2250 if !sent_start {
2251 on_event(AgentEvent::MessageStart {
2252 message: Message::Assistant(Arc::clone(&shared)),
2253 });
2254 sent_start = true;
2255 }
2256 on_event(AgentEvent::MessageUpdate {
2257 message: Message::Assistant(Arc::clone(&shared)),
2258 assistant_message_event: AssistantMessageEvent::TextDelta {
2259 content_index,
2260 delta,
2261 partial: shared,
2262 },
2263 });
2264 }
2265 }
2266 StreamEvent::TextEnd {
2267 content_index,
2268 content,
2269 ..
2270 } => {
2271 self.seed_partial_message_if_missing(&mut added_partial);
2272 if let Some(Message::Assistant(msg_arc)) = self
2273 .messages
2274 .iter_mut()
2275 .rev()
2276 .find(|m| matches!(m, Message::Assistant(_)))
2277 {
2278 {
2279 let msg = Arc::make_mut(msg_arc);
2280 if msg.content.get(content_index).is_none()
2281 && content_index == msg.content.len()
2282 {
2283 msg.content.push(ContentBlock::Text(TextContent::new("")));
2284 }
2285 if let Some(ContentBlock::Text(text)) =
2286 msg.content.get_mut(content_index)
2287 {
2288 text.text.clone_from(&content);
2289 }
2290 }
2291 let shared = Arc::clone(msg_arc);
2292 if !sent_start {
2293 on_event(AgentEvent::MessageStart {
2294 message: Message::Assistant(Arc::clone(&shared)),
2295 });
2296 sent_start = true;
2297 }
2298 on_event(AgentEvent::MessageUpdate {
2299 message: Message::Assistant(Arc::clone(&shared)),
2300 assistant_message_event: AssistantMessageEvent::TextEnd {
2301 content_index,
2302 content,
2303 partial: shared,
2304 },
2305 });
2306 }
2307 }
2308 StreamEvent::ThinkingStart { content_index, .. } => {
2309 self.seed_partial_message_if_missing(&mut added_partial);
2310 if let Some(Message::Assistant(msg_arc)) = self
2311 .messages
2312 .iter_mut()
2313 .rev()
2314 .find(|m| matches!(m, Message::Assistant(_)))
2315 {
2316 let msg = Arc::make_mut(msg_arc);
2317 if content_index == msg.content.len() {
2318 msg.content.push(ContentBlock::Thinking(ThinkingContent {
2319 thinking: String::new(),
2320 thinking_signature: None,
2321 }));
2322 }
2323 let shared = Arc::clone(msg_arc);
2324 if !sent_start {
2325 on_event(AgentEvent::MessageStart {
2326 message: Message::Assistant(Arc::clone(&shared)),
2327 });
2328 sent_start = true;
2329 }
2330 on_event(AgentEvent::MessageUpdate {
2331 message: Message::Assistant(Arc::clone(&shared)),
2332 assistant_message_event: AssistantMessageEvent::ThinkingStart {
2333 content_index,
2334 partial: shared,
2335 },
2336 });
2337 }
2338 }
2339 StreamEvent::ThinkingDelta {
2340 content_index,
2341 delta,
2342 ..
2343 } => {
2344 self.seed_partial_message_if_missing(&mut added_partial);
2345 if let Some(Message::Assistant(msg_arc)) = self
2346 .messages
2347 .iter_mut()
2348 .rev()
2349 .find(|m| matches!(m, Message::Assistant(_)))
2350 {
2351 {
2352 let msg = Arc::make_mut(msg_arc);
2353 if msg.content.get(content_index).is_none()
2354 && content_index == msg.content.len()
2355 {
2356 msg.content.push(ContentBlock::Thinking(ThinkingContent {
2357 thinking: String::new(),
2358 thinking_signature: None,
2359 }));
2360 }
2361 if let Some(ContentBlock::Thinking(thinking)) =
2362 msg.content.get_mut(content_index)
2363 {
2364 thinking.thinking.push_str(&delta);
2365 }
2366 }
2367 let shared = Arc::clone(msg_arc);
2368 if !sent_start {
2369 on_event(AgentEvent::MessageStart {
2370 message: Message::Assistant(Arc::clone(&shared)),
2371 });
2372 sent_start = true;
2373 }
2374 on_event(AgentEvent::MessageUpdate {
2375 message: Message::Assistant(Arc::clone(&shared)),
2376 assistant_message_event: AssistantMessageEvent::ThinkingDelta {
2377 content_index,
2378 delta,
2379 partial: shared,
2380 },
2381 });
2382 }
2383 }
2384 StreamEvent::ThinkingEnd {
2385 content_index,
2386 content,
2387 ..
2388 } => {
2389 self.seed_partial_message_if_missing(&mut added_partial);
2390 if let Some(Message::Assistant(msg_arc)) = self
2391 .messages
2392 .iter_mut()
2393 .rev()
2394 .find(|m| matches!(m, Message::Assistant(_)))
2395 {
2396 {
2397 let msg = Arc::make_mut(msg_arc);
2398 if msg.content.get(content_index).is_none()
2399 && content_index == msg.content.len()
2400 {
2401 msg.content.push(ContentBlock::Thinking(ThinkingContent {
2402 thinking: String::new(),
2403 thinking_signature: None,
2404 }));
2405 }
2406 if let Some(ContentBlock::Thinking(thinking)) =
2407 msg.content.get_mut(content_index)
2408 {
2409 thinking.thinking.clone_from(&content);
2410 }
2411 }
2412 let shared = Arc::clone(msg_arc);
2413 if !sent_start {
2414 on_event(AgentEvent::MessageStart {
2415 message: Message::Assistant(Arc::clone(&shared)),
2416 });
2417 sent_start = true;
2418 }
2419 on_event(AgentEvent::MessageUpdate {
2420 message: Message::Assistant(Arc::clone(&shared)),
2421 assistant_message_event: AssistantMessageEvent::ThinkingEnd {
2422 content_index,
2423 content,
2424 partial: shared,
2425 },
2426 });
2427 }
2428 }
2429 StreamEvent::ToolCallStart {
2430 content_index,
2431 id,
2432 name,
2433 } => {
2434 self.seed_partial_message_if_missing(&mut added_partial);
2435 if let Some(Message::Assistant(msg_arc)) = self
2436 .messages
2437 .iter_mut()
2438 .rev()
2439 .find(|m| matches!(m, Message::Assistant(_)))
2440 {
2441 let msg = Arc::make_mut(msg_arc);
2442 if content_index == msg.content.len() {
2446 msg.content.push(ContentBlock::ToolCall(ToolCall {
2447 id,
2448 name,
2449 arguments: serde_json::Value::Null,
2450 thought_signature: None,
2451 }));
2452 } else if let Some(ContentBlock::ToolCall(tc)) =
2453 msg.content.get_mut(content_index)
2454 {
2455 if tc.id.is_empty() {
2456 tc.id = id;
2457 }
2458 if tc.name.is_empty() {
2459 tc.name = name;
2460 }
2461 }
2462 let shared = Arc::clone(msg_arc);
2463 if !sent_start {
2464 on_event(AgentEvent::MessageStart {
2465 message: Message::Assistant(Arc::clone(&shared)),
2466 });
2467 sent_start = true;
2468 }
2469 on_event(AgentEvent::MessageUpdate {
2470 message: Message::Assistant(Arc::clone(&shared)),
2471 assistant_message_event: AssistantMessageEvent::ToolCallStart {
2472 content_index,
2473 partial: shared,
2474 },
2475 });
2476 }
2477 }
2478 StreamEvent::ToolCallDelta {
2479 content_index,
2480 delta,
2481 ..
2482 } => {
2483 self.seed_partial_message_if_missing(&mut added_partial);
2484 if let Some(Message::Assistant(msg_arc)) = self
2485 .messages
2486 .iter_mut()
2487 .rev()
2488 .find(|m| matches!(m, Message::Assistant(_)))
2489 {
2490 {
2491 let msg = Arc::make_mut(msg_arc);
2492 if msg.content.get(content_index).is_none()
2493 && content_index == msg.content.len()
2494 {
2495 msg.content.push(ContentBlock::ToolCall(ToolCall {
2496 id: String::new(),
2497 name: String::new(),
2498 arguments: serde_json::Value::Null,
2499 thought_signature: None,
2500 }));
2501 }
2502 let raw = tool_call_raw_args.entry(content_index).or_default();
2515 raw.push_str(&delta);
2516 if let Some(partial_args) =
2517 crate::providers::openai::complete_partial_json(raw)
2518 {
2519 if let Some(ContentBlock::ToolCall(tc)) =
2520 msg.content.get_mut(content_index)
2521 {
2522 tc.arguments = partial_args;
2523 }
2524 }
2525 }
2526 let shared = Arc::clone(msg_arc);
2527 if !sent_start {
2528 on_event(AgentEvent::MessageStart {
2529 message: Message::Assistant(Arc::clone(&shared)),
2530 });
2531 sent_start = true;
2532 }
2533 on_event(AgentEvent::MessageUpdate {
2534 message: Message::Assistant(Arc::clone(&shared)),
2535 assistant_message_event: AssistantMessageEvent::ToolCallDelta {
2536 content_index,
2537 delta,
2538 partial: shared,
2539 },
2540 });
2541 }
2542 }
2543 StreamEvent::ToolCallEnd {
2544 content_index,
2545 tool_call,
2546 ..
2547 } => {
2548 self.seed_partial_message_if_missing(&mut added_partial);
2549 if let Some(Message::Assistant(msg_arc)) = self
2550 .messages
2551 .iter_mut()
2552 .rev()
2553 .find(|m| matches!(m, Message::Assistant(_)))
2554 {
2555 {
2556 let msg = Arc::make_mut(msg_arc);
2557 if msg.content.get(content_index).is_none()
2558 && content_index == msg.content.len()
2559 {
2560 msg.content.push(ContentBlock::ToolCall(ToolCall {
2561 id: String::new(),
2562 name: String::new(),
2563 arguments: serde_json::Value::Null,
2564 thought_signature: None,
2565 }));
2566 }
2567 if let Some(ContentBlock::ToolCall(tc)) =
2568 msg.content.get_mut(content_index)
2569 {
2570 *tc = tool_call.clone();
2571 }
2572 }
2573 let shared = Arc::clone(msg_arc);
2574 if !sent_start {
2575 on_event(AgentEvent::MessageStart {
2576 message: Message::Assistant(Arc::clone(&shared)),
2577 });
2578 sent_start = true;
2579 }
2580 on_event(AgentEvent::MessageUpdate {
2581 message: Message::Assistant(Arc::clone(&shared)),
2582 assistant_message_event: AssistantMessageEvent::ToolCallEnd {
2583 content_index,
2584 tool_call,
2585 partial: shared,
2586 },
2587 });
2588 }
2589 }
2590 StreamEvent::Done { message, .. } => {
2591 return Ok(self.finalize_assistant_message(message, &on_event, added_partial));
2592 }
2593 StreamEvent::Error { error, .. } => {
2594 return Ok(self.finalize_assistant_message(error, &on_event, added_partial));
2595 }
2596 }
2597 }
2598
2599 if added_partial {
2603 if let Some(Message::Assistant(last_msg)) = self
2604 .messages
2605 .iter()
2606 .rev()
2607 .find(|m| matches!(m, Message::Assistant(_)))
2608 {
2609 let mut final_msg = (**last_msg).clone();
2610 final_msg.stop_reason = StopReason::Error;
2611 final_msg.error_message = Some("Stream ended without Done event".to_string());
2612 return Ok(self.finalize_assistant_message(final_msg, &on_event, true));
2613 }
2614 }
2615 Err(Error::api("Stream ended without Done event"))
2616 }
2617
2618 fn seed_partial_message_if_missing(&mut self, added_partial: &mut bool) {
2623 if *added_partial {
2624 return;
2625 }
2626
2627 let message = AssistantMessage {
2628 content: Vec::new(),
2629 api: self.provider.api().to_string(),
2630 provider: self.provider.name().to_string(),
2631 model: self.provider.model_id().to_string(),
2632 usage: Usage::default(),
2633 stop_reason: StopReason::Stop,
2634 error_message: None,
2635 timestamp: Utc::now().timestamp_millis(),
2636 };
2637 self.messages.push(Message::Assistant(Arc::new(message)));
2638 *added_partial = true;
2639 }
2640
2641 fn update_partial_message(
2646 &mut self,
2647 partial: Arc<AssistantMessage>,
2648 added_partial: &mut bool,
2649 ) -> bool {
2650 if *added_partial {
2651 if let Some(target) = self
2652 .messages
2653 .iter_mut()
2654 .rev()
2655 .find(|m| matches!(m, Message::Assistant(_)))
2656 {
2657 *target = Message::Assistant(partial);
2658 } else {
2659 tracing::warn!("update_partial_message: expected an Assistant message in history");
2662 self.messages.push(Message::Assistant(partial));
2663 }
2664 false
2665 } else {
2666 self.messages.push(Message::Assistant(partial));
2667 *added_partial = true;
2668 true
2669 }
2670 }
2671
2672 fn finalize_assistant_message(
2673 &mut self,
2674 message: AssistantMessage,
2675 on_event: &Arc<dyn Fn(AgentEvent) + Send + Sync>,
2676 added_partial: bool,
2677 ) -> AssistantMessage {
2678 let arc = Arc::new(message);
2679 if added_partial {
2680 if let Some(target) = self
2681 .messages
2682 .iter_mut()
2683 .rev()
2684 .find(|m| matches!(m, Message::Assistant(_)))
2685 {
2686 *target = Message::Assistant(Arc::clone(&arc));
2687 } else {
2688 tracing::warn!(
2691 "finalize_assistant_message: expected an Assistant message in history"
2692 );
2693 self.messages.push(Message::Assistant(Arc::clone(&arc)));
2694 on_event(AgentEvent::MessageStart {
2695 message: Message::Assistant(Arc::clone(&arc)),
2696 });
2697 }
2698 } else {
2699 self.messages.push(Message::Assistant(Arc::clone(&arc)));
2700 on_event(AgentEvent::MessageStart {
2701 message: Message::Assistant(Arc::clone(&arc)),
2702 });
2703 }
2704
2705 on_event(AgentEvent::MessageEnd {
2706 message: Message::Assistant(Arc::clone(&arc)),
2707 });
2708 Arc::try_unwrap(arc).unwrap_or_else(|a| (*a).clone())
2709 }
2710
2711 async fn execute_tool_batch(
2712 &self,
2713 batch: Vec<(usize, ToolCall)>,
2714 on_event: AgentEventHandler,
2715 abort: Option<AbortSignal>,
2716 latency: SharedTurnLatencyAccumulator,
2717 ) -> Vec<(usize, (ToolOutput, bool))> {
2718 let parallelism = compatible_tool_parallelism_limit();
2719 let futures = batch.into_iter().map(|(idx, tc)| {
2720 let on_event = Arc::clone(&on_event);
2721 let latency = Arc::clone(&latency);
2722 async move { (idx, self.execute_tool_owned(tc, on_event, latency).await) }
2723 });
2724
2725 if let Some(signal) = abort.as_ref() {
2726 use futures::future::{Either, select};
2727 let all_fut = stream::iter(futures)
2728 .buffer_unordered(parallelism)
2729 .collect::<Vec<_>>()
2730 .fuse();
2731 let abort_fut = signal.wait().fuse();
2732 futures::pin_mut!(all_fut, abort_fut);
2733
2734 match select(all_fut, abort_fut).await {
2735 Either::Left((batch_results, _)) => batch_results,
2736 Either::Right(_) => Vec::new(), }
2738 } else {
2739 stream::iter(futures)
2740 .buffer_unordered(parallelism)
2741 .collect::<Vec<_>>()
2742 .await
2743 }
2744 }
2745
2746 #[allow(clippy::too_many_lines)]
2747 async fn execute_tool_calls(
2748 &mut self,
2749 tool_calls: &[ToolCall],
2750 on_event: AgentEventHandler,
2751 new_messages: &mut Vec<Message>,
2752 abort: Option<AbortSignal>,
2753 latency: SharedTurnLatencyAccumulator,
2754 ) -> Result<ToolExecutionOutcome> {
2755 let mut results = Vec::new();
2756 let mut steering_messages: Option<Vec<Message>> = None;
2757
2758 for tool_call in tool_calls {
2760 on_event(AgentEvent::ToolExecutionStart {
2761 tool_call_id: tool_call.id.clone(),
2762 tool_name: tool_call.name.clone(),
2763 args: tool_call.arguments.clone(),
2764 });
2765 }
2766
2767 let effect_plan = tool_calls
2769 .iter()
2770 .map(|tool_call| {
2771 self.tools
2772 .get(&tool_call.name)
2773 .map_or_else(ToolEffects::write, Tool::effects)
2774 })
2775 .collect::<Vec<_>>();
2776 let effect_batches = plan_tool_effect_batches(&effect_plan);
2777 let mut recorded_results: Vec<Option<Arc<ToolResultMessage>>> =
2778 vec![None; tool_calls.len()];
2779
2780 for effect_batch in effect_batches {
2781 if abort.as_ref().is_some_and(AbortSignal::is_aborted) {
2782 break;
2783 }
2784
2785 let steering = self.drain_steering_messages().await;
2786 if !steering.is_empty() {
2787 steering_messages = Some(steering);
2788 break;
2789 }
2790
2791 let batch_len = effect_batch.end.saturating_sub(effect_batch.start);
2792 let batch = tool_calls
2793 .iter()
2794 .cloned()
2795 .enumerate()
2796 .skip(effect_batch.start)
2797 .take(batch_len)
2798 .collect();
2799 let mut batch_results = self
2800 .execute_tool_batch(
2801 batch,
2802 Arc::clone(&on_event),
2803 abort.clone(),
2804 Arc::clone(&latency),
2805 )
2806 .await;
2807 batch_results.sort_by_key(|(idx, _)| *idx);
2808 for (idx, (output, is_error)) in batch_results {
2809 if let (Some(tool_call), Some(recorded_result)) =
2810 (tool_calls.get(idx), recorded_results.get_mut(idx))
2811 {
2812 *recorded_result = Some(self.record_tool_result(
2813 tool_call,
2814 output,
2815 is_error,
2816 &on_event,
2817 new_messages,
2818 ));
2819 }
2820 }
2821 }
2822
2823 for (index, tool_call) in tool_calls.iter().enumerate() {
2825 if steering_messages.is_none() && !abort.as_ref().is_some_and(AbortSignal::is_aborted) {
2828 let steering = self.drain_steering_messages().await;
2829 if !steering.is_empty() {
2830 steering_messages = Some(steering);
2831 }
2832 }
2833
2834 if let Some(tool_result) = recorded_results.get_mut(index).and_then(Option::take) {
2837 results.push(tool_result);
2838 } else if steering_messages.is_some() {
2839 results.push(self.skip_tool_call(tool_call, &on_event, new_messages));
2841 } else {
2842 let output = ToolOutput {
2844 content: vec![ContentBlock::Text(TextContent::new(
2845 "Tool execution aborted",
2846 ))],
2847 details: Some(Self::tool_cancellation_details(
2848 &tool_call.name,
2849 "abort_signal",
2850 )),
2851 is_error: true,
2852 };
2853
2854 on_event(AgentEvent::ToolExecutionUpdate {
2855 tool_call_id: tool_call.id.clone(),
2856 tool_name: tool_call.name.clone(),
2857 args: tool_call.arguments.clone(),
2858 partial_result: ToolOutput {
2859 content: output.content.clone(),
2860 details: output.details.clone(),
2861 is_error: true,
2862 },
2863 });
2864
2865 on_event(AgentEvent::ToolExecutionEnd {
2866 tool_call_id: tool_call.id.clone(),
2867 tool_name: tool_call.name.clone(),
2868 result: ToolOutput {
2869 content: output.content.clone(),
2870 details: output.details.clone(),
2871 is_error: true,
2872 },
2873 is_error: true,
2874 });
2875
2876 let tool_result = Arc::new(ToolResultMessage {
2877 tool_call_id: tool_call.id.clone(),
2878 tool_name: tool_call.name.clone(),
2879 content: output.content,
2880 details: output.details,
2881 is_error: true,
2882 timestamp: Utc::now().timestamp_millis(),
2883 });
2884
2885 let msg = Message::ToolResult(Arc::clone(&tool_result));
2886 self.messages.push(msg.clone());
2887 on_event(AgentEvent::MessageStart {
2888 message: msg.clone(),
2889 });
2890 let end_msg = msg.clone();
2891 new_messages.push(msg);
2892 on_event(AgentEvent::MessageEnd { message: end_msg });
2893
2894 results.push(tool_result);
2895 }
2896 }
2897
2898 Ok(ToolExecutionOutcome {
2899 tool_results: results,
2900 steering_messages,
2901 })
2902 }
2903
2904 fn record_tool_result(
2905 &mut self,
2906 tool_call: &ToolCall,
2907 output: ToolOutput,
2908 is_error: bool,
2909 on_event: &AgentEventHandler,
2910 new_messages: &mut Vec<Message>,
2911 ) -> Arc<ToolResultMessage> {
2912 on_event(AgentEvent::ToolExecutionUpdate {
2913 tool_call_id: tool_call.id.clone(),
2914 tool_name: tool_call.name.clone(),
2915 args: tool_call.arguments.clone(),
2916 partial_result: ToolOutput {
2917 content: output.content.clone(),
2918 details: output.details.clone(),
2919 is_error,
2920 },
2921 });
2922
2923 let tool_result = Arc::new(ToolResultMessage {
2924 tool_call_id: tool_call.id.clone(),
2925 tool_name: tool_call.name.clone(),
2926 content: output.content,
2927 details: output.details,
2928 is_error,
2929 timestamp: Utc::now().timestamp_millis(),
2930 });
2931
2932 on_event(AgentEvent::ToolExecutionEnd {
2933 tool_call_id: tool_result.tool_call_id.clone(),
2934 tool_name: tool_result.tool_name.clone(),
2935 result: ToolOutput {
2936 content: tool_result.content.clone(),
2937 details: tool_result.details.clone(),
2938 is_error,
2939 },
2940 is_error,
2941 });
2942
2943 let msg = Message::ToolResult(Arc::clone(&tool_result));
2944 self.messages.push(msg.clone());
2945 on_event(AgentEvent::MessageStart {
2946 message: msg.clone(),
2947 });
2948 new_messages.push(msg.clone());
2949 on_event(AgentEvent::MessageEnd { message: msg });
2950
2951 tool_result
2952 }
2953
2954 async fn execute_tool(
2955 &self,
2956 tool_call: ToolCall,
2957 on_event: AgentEventHandler,
2958 latency: SharedTurnLatencyAccumulator,
2959 ) -> (ToolOutput, bool) {
2960 let extensions = self.extensions.clone();
2961
2962 let approval_denied_output = self
2963 .request_tool_approval(&tool_call, Arc::clone(&on_event))
2964 .await;
2965
2966 let (mut output, is_error) = if let Some(output) = approval_denied_output {
2967 (output, true)
2968 } else if let Some(extensions) = &extensions {
2969 let hook_started_at = Instant::now();
2970 let hook_outcome = Self::dispatch_tool_call_hook(
2971 extensions,
2972 &tool_call,
2973 self.config.fail_closed_hooks,
2974 )
2975 .await;
2976 record_extension_hostcall_latency(&latency, hook_started_at.elapsed());
2977
2978 if let Some(blocked_output) = hook_outcome {
2979 (blocked_output, true)
2980 } else {
2981 let tool_started_at = Instant::now();
2982 let outcome = self
2983 .execute_tool_without_hooks(&tool_call, Arc::clone(&on_event))
2984 .await;
2985 record_local_tool_latency(&latency, tool_started_at.elapsed());
2986 outcome
2987 }
2988 } else {
2989 let tool_started_at = Instant::now();
2990 let outcome = self
2991 .execute_tool_without_hooks(&tool_call, Arc::clone(&on_event))
2992 .await;
2993 record_local_tool_latency(&latency, tool_started_at.elapsed());
2994 outcome
2995 };
2996
2997 if let Some(extensions) = &extensions {
2998 let hook_started_at = Instant::now();
2999 Self::apply_tool_result_hook(extensions, &tool_call, &mut output, is_error).await;
3000 record_extension_hostcall_latency(&latency, hook_started_at.elapsed());
3001 }
3002
3003 (output, is_error)
3004 }
3005
3006 async fn request_tool_approval(
3007 &self,
3008 tool_call: &ToolCall,
3009 on_event: AgentEventHandler,
3010 ) -> Option<ToolOutput> {
3011 let Some(approval) = &self.config.tool_approval else {
3012 return None;
3013 };
3014
3015 let request = ToolApprovalRequest {
3016 tool_call_id: tool_call.id.clone(),
3017 tool_name: tool_call.name.clone(),
3018 arguments: tool_call.arguments.clone(),
3019 };
3020
3021 match approval(request).await {
3022 ToolApprovalDecision::Allow => {
3023 on_event(AgentEvent::ToolExecutionUpdate {
3024 tool_call_id: tool_call.id.clone(),
3025 tool_name: tool_call.name.clone(),
3026 args: tool_call.arguments.clone(),
3027 partial_result: ToolOutput {
3028 content: Vec::new(),
3029 details: Some(json!({
3030 "schema": TOOL_APPROVAL_STATUS_SCHEMA_V1,
3031 "status": "approved",
3032 })),
3033 is_error: false,
3034 },
3035 });
3036 None
3037 }
3038 ToolApprovalDecision::Deny { reason } => {
3039 Some(Self::tool_approval_denied_output(&reason))
3040 }
3041 }
3042 }
3043
3044 async fn execute_tool_owned(
3045 &self,
3046 tool_call: ToolCall,
3047 on_event: AgentEventHandler,
3048 latency: SharedTurnLatencyAccumulator,
3049 ) -> (ToolOutput, bool) {
3050 self.execute_tool(tool_call, on_event, latency).await
3051 }
3052
3053 async fn execute_tool_without_hooks(
3054 &self,
3055 tool_call: &ToolCall,
3056 on_event: AgentEventHandler,
3057 ) -> (ToolOutput, bool) {
3058 let Some(tool) = self.tools.get(&tool_call.name) else {
3060 return (Self::tool_not_found_output(&tool_call.name), true);
3061 };
3062
3063 let tool_name = tool_call.name.clone();
3064 let tool_id = tool_call.id.clone();
3065 let tool_args = tool_call.arguments.clone();
3066 let on_event = Arc::clone(&on_event);
3067
3068 let update_callback = move |update: ToolUpdate| {
3069 on_event(AgentEvent::ToolExecutionUpdate {
3070 tool_call_id: tool_id.clone(),
3071 tool_name: tool_name.clone(),
3072 args: tool_args.clone(),
3073 partial_result: ToolOutput {
3074 content: update.content,
3075 details: update.details,
3076 is_error: false,
3077 },
3078 });
3079 };
3080
3081 let _artifact_session_guard =
3082 self.config
3083 .stream_options
3084 .session_id
3085 .as_deref()
3086 .map(|session_id| {
3087 crate::tools::register_tool_output_artifact_session(&tool_call.id, session_id)
3088 });
3089
3090 match tool
3091 .execute(
3092 &tool_call.id,
3093 tool_call.arguments.clone(),
3094 Some(Box::new(update_callback)),
3095 )
3096 .await
3097 {
3098 Ok(output) => {
3099 let is_error = output.is_error;
3100 (output, is_error)
3101 }
3102 Err(e) => (
3103 ToolOutput {
3104 content: vec![ContentBlock::Text(TextContent::new(format!("Error: {e}")))],
3105 details: None,
3106 is_error: true,
3107 },
3108 true,
3109 ),
3110 }
3111 }
3112
3113 fn tool_not_found_output(tool_name: &str) -> ToolOutput {
3114 ToolOutput {
3115 content: vec![ContentBlock::Text(TextContent::new(format!(
3116 "Error: Tool '{tool_name}' not found"
3117 )))],
3118 details: None,
3119 is_error: true,
3120 }
3121 }
3122
3123 fn tool_cancellation_details(tool_name: &str, reason: &str) -> Value {
3124 json!({
3125 "schema": TOOL_CANCELLATION_SCHEMA_V1,
3126 "status": "cancelled",
3127 "reason": reason,
3128 "toolName": tool_name,
3129 "cleanup": "tool_result_recorded_no_success",
3130 })
3131 }
3132
3133 async fn dispatch_tool_call_hook(
3134 extensions: &ExtensionManager,
3135 tool_call: &ToolCall,
3136 fail_closed_hooks: bool,
3137 ) -> Option<ToolOutput> {
3138 match extensions
3139 .dispatch_tool_call(tool_call, EXTENSION_EVENT_TIMEOUT_MS)
3140 .await
3141 {
3142 Ok(Some(result)) if result.block => {
3143 Some(Self::tool_call_blocked_output(result.reason.as_deref()))
3144 }
3145 Ok(_) => None,
3146 Err(err) => {
3147 if fail_closed_hooks {
3148 tracing::warn!(
3149 error = ?err,
3150 "tool_call extension hook failed (fail-closed)"
3151 );
3152 Some(Self::tool_call_blocked_output(Some(
3153 "extension hook failed",
3154 )))
3155 } else {
3156 tracing::warn!("tool_call extension hook failed (fail-open): {err}");
3157 None
3158 }
3159 }
3160 }
3161 }
3162
3163 fn tool_call_blocked_output(reason: Option<&str>) -> ToolOutput {
3164 let reason = reason.map(str::trim).filter(|reason| !reason.is_empty());
3165 let message = reason.map_or_else(
3166 || "Tool execution was blocked by an extension".to_string(),
3167 |reason| format!("Tool execution blocked: {reason}"),
3168 );
3169
3170 ToolOutput {
3171 content: vec![ContentBlock::Text(TextContent::new(message))],
3172 details: None,
3173 is_error: true,
3174 }
3175 }
3176
3177 fn tool_approval_denied_output(reason: &str) -> ToolOutput {
3178 let reason = reason.trim();
3179 let reason = if reason.is_empty() {
3180 "tool approval denied"
3181 } else {
3182 reason
3183 };
3184
3185 ToolOutput {
3186 content: vec![ContentBlock::Text(TextContent::new(format!(
3187 "Tool execution denied: {reason}"
3188 )))],
3189 details: Some(json!({
3190 "schema": TOOL_APPROVAL_DENIED_SCHEMA_V1,
3191 "status": "denied",
3192 "reason": reason,
3193 })),
3194 is_error: true,
3195 }
3196 }
3197
3198 async fn apply_tool_result_hook(
3199 extensions: &ExtensionManager,
3200 tool_call: &ToolCall,
3201 output: &mut ToolOutput,
3202 is_error: bool,
3203 ) {
3204 match extensions
3205 .dispatch_tool_result(tool_call, &*output, is_error, EXTENSION_EVENT_TIMEOUT_MS)
3206 .await
3207 {
3208 Ok(Some(result)) => {
3209 if let Some(content) = result.content {
3210 output.content = content;
3211 }
3212 if let Some(details) = result.details {
3213 output.details = Some(details);
3214 }
3215 }
3216 Ok(None) => {}
3217 Err(err) => tracing::warn!("tool_result extension hook failed (fail-open): {err}"),
3218 }
3219 }
3220
3221 fn skip_tool_call(
3222 &mut self,
3223 tool_call: &ToolCall,
3224 on_event: &Arc<dyn Fn(AgentEvent) + Send + Sync>,
3225 new_messages: &mut Vec<Message>,
3226 ) -> Arc<ToolResultMessage> {
3227 let output = ToolOutput {
3228 content: vec![ContentBlock::Text(TextContent::new(
3229 "Skipped due to queued user message.",
3230 ))],
3231 details: None,
3232 is_error: true,
3233 };
3234
3235 on_event(AgentEvent::ToolExecutionUpdate {
3238 tool_call_id: tool_call.id.clone(),
3239 tool_name: tool_call.name.clone(),
3240 args: tool_call.arguments.clone(),
3241 partial_result: output.clone(),
3242 });
3243 on_event(AgentEvent::ToolExecutionEnd {
3244 tool_call_id: tool_call.id.clone(),
3245 tool_name: tool_call.name.clone(),
3246 result: output.clone(),
3247 is_error: true,
3248 });
3249
3250 let tool_result = Arc::new(ToolResultMessage {
3251 tool_call_id: tool_call.id.clone(),
3252 tool_name: tool_call.name.clone(),
3253 content: output.content,
3254 details: output.details,
3255 is_error: true,
3256 timestamp: Utc::now().timestamp_millis(),
3257 });
3258
3259 let msg = Message::ToolResult(Arc::clone(&tool_result));
3260 self.messages.push(msg.clone());
3261 new_messages.push(msg.clone());
3262
3263 on_event(AgentEvent::MessageStart {
3264 message: msg.clone(),
3265 });
3266 on_event(AgentEvent::MessageEnd { message: msg });
3267
3268 tool_result
3269 }
3270}
3271
3272struct ToolExecutionOutcome {
3277 tool_results: Vec<Arc<ToolResultMessage>>,
3278 steering_messages: Option<Vec<Message>>,
3279}
3280
3281pub struct PreWarmedExtensionRuntime {
3286 pub manager: ExtensionManager,
3288 pub runtime: ExtensionRuntimeHandle,
3290 pub tools: Arc<ToolRegistry>,
3292}
3293
3294struct AtomicBoolGuard(Arc<AtomicBool>);
3297
3298impl AtomicBoolGuard {
3299 fn activate(flag: &Arc<AtomicBool>) -> Self {
3300 flag.store(true, Ordering::SeqCst);
3301 Self(Arc::clone(flag))
3302 }
3303}
3304
3305impl Drop for AtomicBoolGuard {
3306 fn drop(&mut self) {
3307 self.0.store(false, Ordering::SeqCst);
3308 }
3309}
3310
3311pub struct AgentSession {
3312 pub agent: Agent,
3313 pub session: Arc<Mutex<Session>>,
3314 save_enabled: bool,
3315 input_source: InputSource,
3316 pub extensions: Option<ExtensionRegion>,
3319 extensions_is_streaming: Arc<AtomicBool>,
3320 extensions_is_compacting: Arc<AtomicBool>,
3321 extensions_turn_active: Arc<AtomicBool>,
3322 extensions_pending_idle_actions: Arc<StdMutex<VecDeque<PendingIdleAction>>>,
3323 extension_queue_modes: Option<Arc<StdMutex<ExtensionQueueModeState>>>,
3324 extension_injected_queue: Option<Arc<StdMutex<ExtensionInjectedQueue>>>,
3325 extension_ai_completion: Arc<StdMutex<ExtensionAiCompletionHostState>>,
3326 compaction_settings: ResolvedCompactionSettings,
3327 compaction_runtime: Option<Runtime>,
3328 runtime_handle: Option<RuntimeHandle>,
3329 compaction_worker: CompactionWorkerState,
3330 model_registry: Option<ModelRegistry>,
3331 auth_storage: Option<AuthStorage>,
3332 api_key_override: Option<String>,
3333 semantic_context_bundle: Option<SemanticContextBundleInjection>,
3334}
3335
3336#[derive(Debug, Clone, Copy)]
3337struct ExtensionQueueModeState {
3338 steering_mode: QueueMode,
3339 follow_up_mode: QueueMode,
3340}
3341
3342impl ExtensionQueueModeState {
3343 const fn new(steering_mode: QueueMode, follow_up_mode: QueueMode) -> Self {
3344 Self {
3345 steering_mode,
3346 follow_up_mode,
3347 }
3348 }
3349
3350 const fn set_modes(&mut self, steering_mode: QueueMode, follow_up_mode: QueueMode) {
3351 self.steering_mode = steering_mode;
3352 self.follow_up_mode = follow_up_mode;
3353 }
3354}
3355
3356#[derive(Debug)]
3357struct ExtensionInjectedQueue {
3358 steering: VecDeque<Message>,
3359 follow_up: VecDeque<Message>,
3360 steering_mode: QueueMode,
3361 follow_up_mode: QueueMode,
3362}
3363
3364impl ExtensionInjectedQueue {
3365 const fn new(steering_mode: QueueMode, follow_up_mode: QueueMode) -> Self {
3366 Self {
3367 steering: VecDeque::new(),
3368 follow_up: VecDeque::new(),
3369 steering_mode,
3370 follow_up_mode,
3371 }
3372 }
3373
3374 const fn set_modes(&mut self, steering_mode: QueueMode, follow_up_mode: QueueMode) {
3375 self.steering_mode = steering_mode;
3376 self.follow_up_mode = follow_up_mode;
3377 }
3378
3379 fn push_steering(&mut self, message: Message) {
3380 if self.steering.len() >= MAX_STEERING_QUEUE_SIZE {
3381 tracing::warn!(
3382 "Extension steering queue full ({} messages), dropping oldest message",
3383 MAX_STEERING_QUEUE_SIZE
3384 );
3385 self.steering.pop_front();
3386 }
3387 self.steering.push_back(message);
3388 }
3389
3390 fn push_follow_up(&mut self, message: Message) {
3391 if self.follow_up.len() >= MAX_FOLLOW_UP_QUEUE_SIZE {
3392 tracing::warn!(
3393 "Extension follow-up queue full ({} messages), dropping oldest message",
3394 MAX_FOLLOW_UP_QUEUE_SIZE
3395 );
3396 self.follow_up.pop_front();
3397 }
3398 self.follow_up.push_back(message);
3399 }
3400
3401 fn pop_steering(&mut self) -> Vec<Message> {
3402 match self.steering_mode {
3403 QueueMode::All => self.steering.drain(..).collect(),
3404 QueueMode::OneAtATime => self.steering.pop_front().into_iter().collect(),
3405 }
3406 }
3407
3408 fn pop_follow_up(&mut self) -> Vec<Message> {
3409 match self.follow_up_mode {
3410 QueueMode::All => self.follow_up.drain(..).collect(),
3411 QueueMode::OneAtATime => self.follow_up.pop_front().into_iter().collect(),
3412 }
3413 }
3414}
3415
3416impl Default for ExtensionInjectedQueue {
3417 fn default() -> Self {
3418 Self::new(QueueMode::OneAtATime, QueueMode::OneAtATime)
3419 }
3420}
3421
3422#[derive(Debug)]
3423enum PendingIdleAction {
3424 CustomMessage(Message),
3425 UserText(String),
3426}
3427
3428#[derive(Clone)]
3429struct AgentSessionHostActions {
3430 session: Arc<Mutex<Session>>,
3431 injected: Arc<StdMutex<ExtensionInjectedQueue>>,
3432 is_streaming: Arc<AtomicBool>,
3433 is_turn_active: Arc<AtomicBool>,
3434 pending_idle_actions: Arc<StdMutex<VecDeque<PendingIdleAction>>>,
3435 ai_completion: Arc<StdMutex<ExtensionAiCompletionHostState>>,
3436}
3437
3438#[derive(Clone)]
3439struct ExtensionAiCompletionHostState {
3440 provider: Arc<dyn Provider>,
3441 stream_options: StreamOptions,
3442 models: Vec<Value>,
3443}
3444
3445impl AgentSessionHostActions {
3446 fn enqueue(&self, deliver_as: Option<ExtensionDeliverAs>, message: Message) {
3447 let deliver_as = deliver_as.unwrap_or(ExtensionDeliverAs::Steer);
3448 let Ok(mut queue) = self.injected.lock() else {
3449 tracing::error!("injected queue mutex poisoned; dropping extension message");
3450 return;
3451 };
3452 match deliver_as {
3453 ExtensionDeliverAs::FollowUp => {
3454 queue.push_follow_up(message);
3455 }
3456 ExtensionDeliverAs::Steer | ExtensionDeliverAs::NextTurn => {
3457 queue.push_steering(message);
3458 }
3459 }
3460 }
3461
3462 async fn append_to_session(&self, message: Message) -> Result<()> {
3463 let cx = crate::agent_cx::AgentCx::for_current_or_request();
3464 let mut session = self
3465 .session
3466 .lock(cx.cx())
3467 .await
3468 .map_err(|e| Error::session(e.to_string()))?;
3469 session.append_model_message(message);
3470 Ok(())
3471 }
3472
3473 fn queue_pending_idle_action(&self, action: PendingIdleAction) {
3474 let Ok(mut actions) = self.pending_idle_actions.lock() else {
3475 tracing::error!("pending idle actions mutex poisoned; dropping idle action");
3476 return;
3477 };
3478 actions.push_back(action);
3479 }
3480}
3481
3482#[async_trait]
3483impl ExtensionHostActions for AgentSessionHostActions {
3484 async fn send_message(&self, message: ExtensionSendMessage) -> Result<()> {
3485 let custom_message = Message::Custom(CustomMessage {
3486 content: message.content,
3487 custom_type: message.custom_type,
3488 display: message.display,
3489 details: message.details,
3490 timestamp: Utc::now().timestamp_millis(),
3491 });
3492
3493 if matches!(message.deliver_as, Some(ExtensionDeliverAs::NextTurn)) {
3494 return self.append_to_session(custom_message).await;
3495 }
3496
3497 if self.is_streaming.load(Ordering::SeqCst) {
3498 self.enqueue(message.deliver_as, custom_message);
3499 return Ok(());
3500 }
3501
3502 if self.is_turn_active.load(Ordering::SeqCst) {
3503 return self.append_to_session(custom_message).await;
3504 }
3505
3506 if message.trigger_turn {
3507 self.queue_pending_idle_action(PendingIdleAction::CustomMessage(custom_message));
3508 return Ok(());
3509 }
3510
3511 self.append_to_session(custom_message).await
3512 }
3513
3514 async fn send_user_message(&self, message: ExtensionSendUserMessage) -> Result<()> {
3515 let text = message.text;
3516 let user_message = Message::User(UserMessage {
3517 content: UserContent::Text(text.clone()),
3518 timestamp: Utc::now().timestamp_millis(),
3519 });
3520
3521 if self.is_streaming.load(Ordering::SeqCst) {
3522 self.enqueue(message.deliver_as, user_message);
3523 return Ok(());
3524 }
3525
3526 if self.is_turn_active.load(Ordering::SeqCst) {
3527 return self.append_to_session(user_message).await;
3528 }
3529
3530 self.queue_pending_idle_action(PendingIdleAction::UserText(text));
3531 Ok(())
3532 }
3533
3534 async fn complete_ai(&self, request: ExtensionAiCompletionRequest) -> Result<Value> {
3535 let (provider, mut stream_options) = {
3536 let state = self.ai_completion.lock().map_err(|_| {
3537 Error::extension("extension completion host state mutex poisoned".to_string())
3538 })?;
3539 (Arc::clone(&state.provider), state.stream_options.clone())
3540 };
3541
3542 apply_pi_ai_completion_options(&request.options, &mut stream_options)?;
3543 let context = build_pi_ai_completion_context(&request)?;
3544 let provider_name = provider.name().to_string();
3545 let mut events = provider.stream(&context, &stream_options).await?;
3546 let mut streamed_text = String::new();
3547
3548 while let Some(event) = events.next().await {
3549 match event.map_err(|err| Error::provider(provider_name.clone(), err.to_string()))? {
3550 StreamEvent::TextDelta { delta, .. } => streamed_text.push_str(&delta),
3551 StreamEvent::TextEnd { content, .. } => {
3552 streamed_text.push_str(&content);
3553 }
3554 StreamEvent::Done { message, .. } => {
3555 if message.stop_reason == StopReason::Error {
3556 return Err(Error::provider(
3557 provider_name,
3558 pi_ai_assistant_error_message(&message),
3559 ));
3560 }
3561 return pi_ai_completion_response(&message, request.simple);
3562 }
3563 StreamEvent::Error { error, .. } => {
3564 return Err(Error::provider(
3565 provider_name,
3566 pi_ai_assistant_error_message(&error),
3567 ));
3568 }
3569 StreamEvent::Start { .. }
3570 | StreamEvent::TextStart { .. }
3571 | StreamEvent::ThinkingStart { .. }
3572 | StreamEvent::ThinkingDelta { .. }
3573 | StreamEvent::ThinkingEnd { .. }
3574 | StreamEvent::ToolCallStart { .. }
3575 | StreamEvent::ToolCallDelta { .. }
3576 | StreamEvent::ToolCallEnd { .. } => {}
3577 }
3578 }
3579
3580 let suffix = if streamed_text.is_empty() {
3581 String::new()
3582 } else {
3583 format!(" after streaming {} text bytes", streamed_text.len())
3584 };
3585 Err(Error::provider(
3586 provider_name,
3587 format!("pi-ai completion stream ended without Done event{suffix}"),
3588 ))
3589 }
3590
3591 async fn list_ai_models(&self) -> Result<Value> {
3592 let state = self.ai_completion.lock().map_err(|_| {
3593 Error::extension("extension completion host state mutex poisoned".to_string())
3594 })?;
3595 if state.models.is_empty() {
3596 return Ok(json!([{
3597 "id": state.provider.model_id(),
3598 "name": state.provider.model_id(),
3599 "api": state.provider.api(),
3600 "provider": state.provider.name(),
3601 }]));
3602 }
3603 Ok(Value::Array(state.models.clone()))
3604 }
3605}
3606
3607fn pi_ai_model_entry_value(entry: &ModelEntry) -> Value {
3608 json!({
3609 "id": entry.model.id,
3610 "name": entry.model.name,
3611 "api": entry.model.api,
3612 "provider": entry.model.provider,
3613 "baseUrl": entry.model.base_url,
3614 "reasoning": entry.model.reasoning,
3615 "input": entry.model.input,
3616 "cost": entry.model.cost,
3617 "contextWindow": entry.model.context_window,
3618 "maxTokens": entry.model.max_tokens,
3619 "authHeader": entry.auth_header,
3620 "hasCredentials": entry.api_key.is_some(),
3621 })
3622}
3623
3624fn pi_ai_model_registry_values(registry: &ModelRegistry) -> Vec<Value> {
3625 registry
3626 .models()
3627 .iter()
3628 .map(pi_ai_model_entry_value)
3629 .collect()
3630}
3631
3632fn apply_pi_ai_completion_options(
3633 options: &Value,
3634 stream_options: &mut StreamOptions,
3635) -> Result<()> {
3636 if let Some(value) = options
3637 .get("temperature")
3638 .or_else(|| options.get("temp"))
3639 .filter(|value| !value.is_null())
3640 {
3641 let temperature = serde_json::from_value::<f32>(value.clone()).map_err(|err| {
3642 Error::validation(format!(
3643 "pi-ai completion temperature must be numeric: {err}"
3644 ))
3645 })?;
3646 if !(0.0..=2.0).contains(&temperature) {
3647 return Err(Error::validation(
3648 "pi-ai completion temperature must be between 0 and 2".to_string(),
3649 ));
3650 }
3651 stream_options.temperature = Some(temperature);
3652 }
3653
3654 if let Some(value) = options
3655 .get("maxTokens")
3656 .or_else(|| options.get("max_tokens"))
3657 .filter(|value| !value.is_null())
3658 {
3659 let raw = value.as_u64().ok_or_else(|| {
3660 Error::validation("pi-ai completion maxTokens must be an unsigned integer".to_string())
3661 })?;
3662 let max_tokens = u32::try_from(raw).map_err(|_| {
3663 Error::validation("pi-ai completion maxTokens exceeds u32::MAX".to_string())
3664 })?;
3665 if max_tokens == 0 {
3666 return Err(Error::validation(
3667 "pi-ai completion maxTokens must be greater than zero".to_string(),
3668 ));
3669 }
3670 stream_options.max_tokens = Some(max_tokens);
3671 }
3672
3673 Ok(())
3674}
3675
3676fn build_pi_ai_completion_context(
3677 request: &ExtensionAiCompletionRequest,
3678) -> Result<Context<'static>> {
3679 let mut system_prompts = Vec::new();
3680 let mut messages = Vec::new();
3681 collect_pi_ai_context_messages(&request.context, &mut system_prompts, &mut messages)?;
3682
3683 if messages.is_empty() {
3684 return Err(Error::validation(
3685 "@mariozechner/pi-ai completion requires at least one user or assistant message"
3686 .to_string(),
3687 ));
3688 }
3689
3690 let system_prompt = system_prompts
3691 .into_iter()
3692 .filter(|text| !text.trim().is_empty())
3693 .collect::<Vec<_>>()
3694 .join("\n\n");
3695 Ok(Context::owned(
3696 if system_prompt.is_empty() {
3697 None
3698 } else {
3699 Some(system_prompt)
3700 },
3701 messages,
3702 Vec::new(),
3703 ))
3704}
3705
3706fn collect_pi_ai_context_messages(
3707 value: &Value,
3708 system_prompts: &mut Vec<String>,
3709 messages: &mut Vec<Message>,
3710) -> Result<()> {
3711 match value {
3712 Value::Null => {}
3713 Value::String(text) => push_pi_ai_user_message(text, messages),
3714 Value::Array(items) => {
3715 for item in items {
3716 push_pi_ai_message(item, system_prompts, messages)?;
3717 }
3718 }
3719 Value::Object(map) => {
3720 if let Some(system) = map
3721 .get("systemPrompt")
3722 .or_else(|| map.get("system_prompt"))
3723 .or_else(|| map.get("system"))
3724 .and_then(pi_ai_text_from_value)
3725 {
3726 system_prompts.push(system);
3727 }
3728
3729 if let Some(items) = map.get("messages").and_then(Value::as_array) {
3730 for item in items {
3731 push_pi_ai_message(item, system_prompts, messages)?;
3732 }
3733 } else if let Some(prompt) = map
3734 .get("prompt")
3735 .or_else(|| map.get("input"))
3736 .or_else(|| map.get("message"))
3737 .and_then(pi_ai_text_from_value)
3738 {
3739 push_pi_ai_user_message(&prompt, messages);
3740 } else if map.contains_key("role") {
3741 push_pi_ai_message(value, system_prompts, messages)?;
3742 }
3743 }
3744 Value::Bool(_) | Value::Number(_) => push_pi_ai_user_message(&value.to_string(), messages),
3745 }
3746 Ok(())
3747}
3748
3749fn push_pi_ai_message(
3750 value: &Value,
3751 system_prompts: &mut Vec<String>,
3752 messages: &mut Vec<Message>,
3753) -> Result<()> {
3754 let Value::Object(map) = value else {
3755 if let Some(text) = pi_ai_text_from_value(value) {
3756 push_pi_ai_user_message(&text, messages);
3757 }
3758 return Ok(());
3759 };
3760
3761 let role = map
3762 .get("role")
3763 .and_then(Value::as_str)
3764 .unwrap_or("user")
3765 .trim()
3766 .to_ascii_lowercase();
3767 let content = map
3768 .get("content")
3769 .or_else(|| map.get("text"))
3770 .and_then(pi_ai_text_from_value)
3771 .unwrap_or_default();
3772
3773 match role.as_str() {
3774 "system" => {
3775 if !content.trim().is_empty() {
3776 system_prompts.push(content);
3777 }
3778 }
3779 "user" => push_pi_ai_user_message(&content, messages),
3780 "assistant" => push_pi_ai_assistant_message(&content, messages),
3781 other => {
3782 return Err(Error::validation(format!(
3783 "@mariozechner/pi-ai completion does not support {other:?} context messages"
3784 )));
3785 }
3786 }
3787 Ok(())
3788}
3789
3790fn push_pi_ai_user_message(text: &str, messages: &mut Vec<Message>) {
3791 messages.push(Message::User(UserMessage {
3792 content: UserContent::Text(text.to_string()),
3793 timestamp: Utc::now().timestamp_millis(),
3794 }));
3795}
3796
3797fn push_pi_ai_assistant_message(text: &str, messages: &mut Vec<Message>) {
3798 messages.push(Message::assistant(AssistantMessage {
3799 content: vec![ContentBlock::Text(TextContent::new(text.to_string()))],
3800 timestamp: Utc::now().timestamp_millis(),
3801 ..AssistantMessage::default()
3802 }));
3803}
3804
3805fn pi_ai_text_from_value(value: &Value) -> Option<String> {
3806 match value {
3807 Value::Null => None,
3808 Value::String(text) => Some(text.clone()),
3809 Value::Bool(_) | Value::Number(_) => Some(value.to_string()),
3810 Value::Array(items) => {
3811 let mut text = String::new();
3812 for item in items {
3813 if let Some(part) = pi_ai_text_from_value(item)
3814 && !part.is_empty()
3815 {
3816 text.push_str(&part);
3817 }
3818 }
3819 Some(text)
3820 }
3821 Value::Object(map) => map
3822 .get("text")
3823 .or_else(|| map.get("content"))
3824 .or_else(|| map.get("delta"))
3825 .and_then(pi_ai_text_from_value),
3826 }
3827}
3828
3829fn pi_ai_assistant_text(message: &AssistantMessage) -> String {
3830 let mut text = String::new();
3831 for block in &message.content {
3832 if let ContentBlock::Text(text_block) = block {
3833 text.push_str(&text_block.text);
3834 }
3835 }
3836 text
3837}
3838
3839fn pi_ai_assistant_error_message(message: &AssistantMessage) -> String {
3840 message
3841 .error_message
3842 .clone()
3843 .filter(|text| !text.trim().is_empty())
3844 .unwrap_or_else(|| {
3845 let text = pi_ai_assistant_text(message);
3846 if text.trim().is_empty() {
3847 "provider returned an error without a message".to_string()
3848 } else {
3849 text
3850 }
3851 })
3852}
3853
3854fn pi_ai_completion_response(message: &AssistantMessage, simple: bool) -> Result<Value> {
3855 let text = pi_ai_assistant_text(message);
3856 if simple {
3857 return Ok(Value::String(text));
3858 }
3859
3860 Ok(json!({
3861 "message": serde_json::to_value(message)?,
3862 "content": serde_json::to_value(&message.content)?,
3863 "text": text,
3864 "usage": serde_json::to_value(&message.usage)?,
3865 "model": message.model,
3866 "provider": message.provider,
3867 "api": message.api,
3868 "stopReason": message.stop_reason,
3869 }))
3870}
3871
3872#[cfg(test)]
3873mod message_queue_tests {
3874 use super::*;
3875
3876 fn user_message(text: &str) -> Message {
3877 Message::User(UserMessage {
3878 content: UserContent::Text(text.to_string()),
3879 timestamp: 0,
3880 })
3881 }
3882
3883 #[test]
3884 fn message_queue_one_at_a_time() {
3885 let mut queue = MessageQueue::new(QueueMode::OneAtATime, QueueMode::OneAtATime);
3886 queue.push_steering(user_message("a"));
3887 queue.push_steering(user_message("b"));
3888
3889 let first = queue.pop_steering();
3890 assert_eq!(first.len(), 1);
3891 assert!(matches!(
3892 first.first(),
3893 Some(Message::User(UserMessage { content, .. }))
3894 if matches!(content, UserContent::Text(text) if text == "a")
3895 ));
3896
3897 let second = queue.pop_steering();
3898 assert_eq!(second.len(), 1);
3899 assert!(matches!(
3900 second.first(),
3901 Some(Message::User(UserMessage { content, .. }))
3902 if matches!(content, UserContent::Text(text) if text == "b")
3903 ));
3904
3905 assert!(queue.pop_steering().is_empty());
3906 }
3907
3908 #[test]
3909 fn message_queue_all_mode() {
3910 let mut queue = MessageQueue::new(QueueMode::All, QueueMode::OneAtATime);
3911 queue.push_steering(user_message("a"));
3912 queue.push_steering(user_message("b"));
3913
3914 let drained = queue.pop_steering();
3915 assert_eq!(drained.len(), 2);
3916 assert!(queue.pop_steering().is_empty());
3917 }
3918
3919 #[test]
3920 fn message_queue_separates_kinds() {
3921 let mut queue = MessageQueue::new(QueueMode::OneAtATime, QueueMode::OneAtATime);
3922 queue.push_steering(user_message("steer"));
3923 queue.push_follow_up(user_message("follow"));
3924
3925 let steering = queue.pop_steering();
3926 assert_eq!(steering.len(), 1);
3927 assert_eq!(queue.pending_count(), 1);
3928
3929 let follow = queue.pop_follow_up();
3930 assert_eq!(follow.len(), 1);
3931 assert_eq!(queue.pending_count(), 0);
3932 }
3933
3934 #[test]
3935 fn message_queue_seq_increments() {
3936 let mut queue = MessageQueue::new(QueueMode::OneAtATime, QueueMode::OneAtATime);
3937 let first = queue.push_steering(user_message("a"));
3938 let second = queue.push_follow_up(user_message("b"));
3939 assert!(second > first);
3940 }
3941
3942 #[test]
3943 fn message_queue_seq_saturates_at_u64_max() {
3944 let mut queue = MessageQueue::new(QueueMode::OneAtATime, QueueMode::OneAtATime);
3945 queue.next_seq = u64::MAX;
3946
3947 let first = queue.push_steering(user_message("a"));
3948 let second = queue.push_follow_up(user_message("b"));
3949
3950 assert_eq!(first, u64::MAX);
3951 assert_eq!(second, u64::MAX);
3952 assert_eq!(queue.pending_count(), 2);
3953 }
3954
3955 #[test]
3956 fn message_queue_follow_up_all_mode_drains_entire_queue_in_order() {
3957 let mut queue = MessageQueue::new(QueueMode::OneAtATime, QueueMode::All);
3958 queue.push_follow_up(user_message("f1"));
3959 queue.push_follow_up(user_message("f2"));
3960
3961 let follow_up = queue.pop_follow_up();
3962 assert_eq!(follow_up.len(), 2);
3963 assert!(matches!(
3964 follow_up.first(),
3965 Some(Message::User(UserMessage { content, .. }))
3966 if matches!(content, UserContent::Text(text) if text == "f1")
3967 ));
3968 assert!(matches!(
3969 follow_up.get(1),
3970 Some(Message::User(UserMessage { content, .. }))
3971 if matches!(content, UserContent::Text(text) if text == "f2")
3972 ));
3973 assert!(queue.pop_follow_up().is_empty());
3974 }
3975}
3976
3977#[cfg(test)]
3978mod compatible_tool_parallelism_tests {
3979 use super::*;
3980
3981 #[test]
3982 fn compatible_tool_parallelism_preserves_historical_floor() {
3983 assert_eq!(resolve_compatible_tool_parallelism(None, 1), 8);
3984 assert_eq!(resolve_compatible_tool_parallelism(None, 8), 8);
3985 }
3986
3987 #[test]
3988 fn compatible_tool_parallelism_scales_on_many_core_hosts() {
3989 assert_eq!(resolve_compatible_tool_parallelism(None, 32), 32);
3990 assert_eq!(resolve_compatible_tool_parallelism(None, 64), 64);
3991 assert_eq!(resolve_compatible_tool_parallelism(None, 128), 64);
3992 }
3993
3994 #[test]
3995 fn compatible_tool_parallelism_accepts_bounded_override() {
3996 assert_eq!(resolve_compatible_tool_parallelism(Some("16"), 4), 16);
3997 assert_eq!(resolve_compatible_tool_parallelism(Some("512"), 64), 256);
3998 assert_eq!(resolve_compatible_tool_parallelism(Some("1"), 64), 1);
3999 }
4000
4001 #[test]
4002 fn compatible_tool_parallelism_ignores_invalid_override() {
4003 assert_eq!(
4004 resolve_compatible_tool_parallelism(Some("not-a-number"), 24),
4005 24
4006 );
4007 assert_eq!(resolve_compatible_tool_parallelism(Some("0"), 24), 24);
4008 assert_eq!(resolve_compatible_tool_parallelism(Some(" "), 24), 24);
4009 }
4010}
4011
4012#[cfg(test)]
4013mod tool_effect_batch_planning_tests {
4014 use super::*;
4015
4016 #[derive(Debug, Clone, Copy)]
4017 enum SyntheticOutcome {
4018 Success,
4019 Error,
4020 }
4021
4022 #[derive(Debug, Clone)]
4023 struct SyntheticToolCase {
4024 id: String,
4025 name: String,
4026 registered_effects: Option<ToolEffects>,
4027 outcome: SyntheticOutcome,
4028 }
4029
4030 #[derive(Debug, Clone, Copy)]
4031 enum BatchArrivalOrder {
4032 Forward,
4033 Reverse,
4034 RotateLeft(usize),
4035 }
4036
4037 #[derive(Debug, Clone, PartialEq, Eq)]
4038 struct TranscriptEntry {
4039 tool_call_id: String,
4040 tool_name: String,
4041 text: String,
4042 details: serde_json::Value,
4043 is_error: bool,
4044 }
4045
4046 fn batch_ranges(effects: &[ToolEffects]) -> Vec<(usize, usize)> {
4047 plan_tool_effect_batches(effects)
4048 .into_iter()
4049 .map(|batch| (batch.start, batch.end))
4050 .collect()
4051 }
4052
4053 fn batch_plan_json(effects: &[ToolEffects], parallelism_cap: usize) -> serde_json::Value {
4054 serde_json::to_value(tool_effect_batch_plan_evidence(effects, parallelism_cap))
4055 .expect("tool-effect batch evidence should serialize")
4056 }
4057
4058 fn synthetic_tool_case(
4059 index: usize,
4060 name: impl Into<String>,
4061 registered_effects: Option<ToolEffects>,
4062 outcome: SyntheticOutcome,
4063 ) -> SyntheticToolCase {
4064 SyntheticToolCase {
4065 id: format!("call-{index:03}"),
4066 name: name.into(),
4067 registered_effects,
4068 outcome,
4069 }
4070 }
4071
4072 fn effect_plan(cases: &[SyntheticToolCase]) -> Vec<ToolEffects> {
4073 cases
4074 .iter()
4075 .map(|case| case.registered_effects.unwrap_or_else(ToolEffects::write))
4076 .collect()
4077 }
4078
4079 fn make_tool_result(case: &SyntheticToolCase, index: usize) -> ToolResultMessage {
4080 let (content, is_error) = match case.outcome {
4081 SyntheticOutcome::Success => (format!("ok:{}", case.name), false),
4082 SyntheticOutcome::Error => (format!("error:{}", case.name), true),
4083 };
4084 ToolResultMessage {
4085 tool_call_id: case.id.clone(),
4086 tool_name: case.name.clone(),
4087 content: vec![ContentBlock::Text(TextContent::new(content))],
4088 details: Some(serde_json::json!({
4089 "ordinal": index,
4090 "tool": case.name,
4091 "status": if is_error { "error" } else { "ok" },
4092 })),
4093 is_error,
4094 timestamp: 42,
4095 }
4096 }
4097
4098 fn transcript_entry(message: &ToolResultMessage) -> TranscriptEntry {
4099 assert_eq!(message.content.len(), 1, "synthetic result content drifted");
4100 let text = message
4101 .content
4102 .first()
4103 .and_then(|block| match block {
4104 ContentBlock::Text(text) => Some(text.text.clone()),
4105 _ => None,
4106 })
4107 .unwrap_or_else(|| "non-text synthetic result".to_string());
4108 TranscriptEntry {
4109 tool_call_id: message.tool_call_id.clone(),
4110 tool_name: message.tool_name.clone(),
4111 text,
4112 details: message.details.clone().unwrap_or(serde_json::Value::Null),
4113 is_error: message.is_error,
4114 }
4115 }
4116
4117 fn sequential_oracle(cases: &[SyntheticToolCase]) -> Vec<TranscriptEntry> {
4118 cases
4119 .iter()
4120 .enumerate()
4121 .map(|(index, case)| transcript_entry(&make_tool_result(case, index)))
4122 .collect()
4123 }
4124
4125 fn reorder_batch(indices: &mut [usize], order: BatchArrivalOrder) {
4126 match order {
4127 BatchArrivalOrder::Forward => {}
4128 BatchArrivalOrder::Reverse => indices.reverse(),
4129 BatchArrivalOrder::RotateLeft(amount) => {
4130 if !indices.is_empty() {
4131 indices.rotate_left(amount % indices.len());
4132 }
4133 }
4134 }
4135 }
4136
4137 fn scheduled_transcript(
4138 cases: &[SyntheticToolCase],
4139 order: BatchArrivalOrder,
4140 ) -> Vec<TranscriptEntry> {
4141 let effects = effect_plan(cases);
4142 let batches = plan_tool_effect_batches(&effects);
4143 let mut recorded_results: Vec<Option<ToolResultMessage>> = vec![None; cases.len()];
4144
4145 for batch in batches {
4146 let mut completion_order = (batch.start..batch.end).collect::<Vec<_>>();
4147 reorder_batch(&mut completion_order, order);
4148 let mut batch_results = completion_order
4149 .into_iter()
4150 .filter_map(|index| {
4151 cases
4152 .get(index)
4153 .map(|case| (index, make_tool_result(case, index)))
4154 })
4155 .collect::<Vec<_>>();
4156 batch_results.sort_by_key(|(index, _)| *index);
4157 for (index, result) in batch_results {
4158 if let Some(slot) = recorded_results.get_mut(index) {
4159 *slot = Some(result);
4160 }
4161 }
4162 }
4163
4164 assert!(
4165 recorded_results.iter().all(Option::is_some),
4166 "scheduled execution should record every result"
4167 );
4168 recorded_results
4169 .into_iter()
4170 .flatten()
4171 .map(|result| transcript_entry(&result))
4172 .collect()
4173 }
4174
4175 fn assert_barrier_effects_are_singleton_batches(cases: &[SyntheticToolCase]) {
4176 let effects = effect_plan(cases);
4177 for batch in plan_tool_effect_batches(&effects) {
4178 let batch_effects = effects
4179 .get(batch.start..batch.end)
4180 .unwrap_or(&[])
4181 .iter()
4182 .copied()
4183 .fold(ToolEffects::read(), ToolEffects::union);
4184 if !batch_effects.parallel_safe() {
4185 assert_eq!(
4186 batch.end - batch.start,
4187 1,
4188 "barrier batch must serialize original index {}",
4189 batch.start
4190 );
4191 }
4192 }
4193 }
4194
4195 #[test]
4196 fn read_and_network_effects_share_compatible_batch() {
4197 let ranges = batch_ranges(&[
4198 ToolEffects::read(),
4199 ToolEffects::network(),
4200 ToolEffects::read(),
4201 ]);
4202
4203 assert_eq!(ranges, vec![(0, 3)]);
4204 }
4205
4206 #[test]
4207 fn evidence_records_64_plus_compatible_batch_with_parallelism_cap() {
4208 let effects = (0..72)
4209 .map(|index| {
4210 if index % 3 == 0 {
4211 ToolEffects::network()
4212 } else {
4213 ToolEffects::read()
4214 }
4215 })
4216 .collect::<Vec<_>>();
4217
4218 assert_eq!(
4219 batch_plan_json(&effects, 64),
4220 serde_json::json!({
4221 "schema": TOOL_EFFECT_BATCH_PLAN_SCHEMA_V1,
4222 "toolCount": 72,
4223 "parallelismCap": 64,
4224 "batches": [
4225 {
4226 "start": 0,
4227 "end": 72,
4228 "len": 72,
4229 "combinedEffects": ["read", "network"],
4230 "parallelSafe": true
4231 }
4232 ]
4233 })
4234 );
4235 }
4236
4237 #[test]
4238 fn write_effect_creates_deterministic_barrier() {
4239 let ranges = batch_ranges(&[
4240 ToolEffects::read(),
4241 ToolEffects::read(),
4242 ToolEffects::write(),
4243 ToolEffects::read(),
4244 ]);
4245
4246 assert_eq!(ranges, vec![(0, 2), (2, 3), (3, 4)]);
4247 }
4248
4249 #[test]
4250 fn append_and_process_effects_remain_serialized() {
4251 let ranges = batch_ranges(&[
4252 ToolEffects::append(),
4253 ToolEffects::append(),
4254 ToolEffects::process(),
4255 ToolEffects::read(),
4256 ]);
4257
4258 assert_eq!(ranges, vec![(0, 1), (1, 2), (2, 3), (3, 4)]);
4259 }
4260
4261 #[test]
4262 fn combined_process_write_effect_is_exclusive() {
4263 let ranges = batch_ranges(&[
4264 ToolEffects::read(),
4265 ToolEffects::process().union(ToolEffects::write()),
4266 ToolEffects::network(),
4267 ]);
4268
4269 assert_eq!(ranges, vec![(0, 1), (1, 2), (2, 3)]);
4270 }
4271
4272 #[test]
4273 fn evidence_records_barrier_reasons_for_mixed_effects() {
4274 let effects = [
4275 ToolEffects::read(),
4276 ToolEffects::network(),
4277 ToolEffects::write(),
4278 ToolEffects::append(),
4279 ToolEffects::process(),
4280 ToolEffects::read(),
4281 ToolEffects::process().union(ToolEffects::write()),
4282 ];
4283
4284 assert_eq!(
4285 batch_plan_json(&effects, 32),
4286 serde_json::json!({
4287 "schema": TOOL_EFFECT_BATCH_PLAN_SCHEMA_V1,
4288 "toolCount": 7,
4289 "parallelismCap": 32,
4290 "batches": [
4291 {
4292 "start": 0,
4293 "end": 2,
4294 "len": 2,
4295 "combinedEffects": ["read", "network"],
4296 "parallelSafe": true
4297 },
4298 {
4299 "start": 2,
4300 "end": 3,
4301 "len": 1,
4302 "combinedEffects": ["write"],
4303 "parallelSafe": false,
4304 "barrierReason": "write_barrier"
4305 },
4306 {
4307 "start": 3,
4308 "end": 4,
4309 "len": 1,
4310 "combinedEffects": ["append"],
4311 "parallelSafe": false,
4312 "barrierReason": "append_barrier"
4313 },
4314 {
4315 "start": 4,
4316 "end": 5,
4317 "len": 1,
4318 "combinedEffects": ["process"],
4319 "parallelSafe": false,
4320 "barrierReason": "process_barrier"
4321 },
4322 {
4323 "start": 5,
4324 "end": 6,
4325 "len": 1,
4326 "combinedEffects": ["read"],
4327 "parallelSafe": true
4328 },
4329 {
4330 "start": 6,
4331 "end": 7,
4332 "len": 1,
4333 "combinedEffects": ["write", "process"],
4334 "parallelSafe": false,
4335 "barrierReason": "write_process_barrier"
4336 }
4337 ]
4338 })
4339 );
4340 }
4341
4342 #[test]
4343 fn metamorphic_empty_tool_batch_matches_sequential_oracle() {
4344 let cases = Vec::new();
4345
4346 assert!(plan_tool_effect_batches(&effect_plan(&cases)).is_empty());
4347 assert_eq!(
4348 scheduled_transcript(&cases, BatchArrivalOrder::Forward),
4349 sequential_oracle(&cases)
4350 );
4351 }
4352
4353 #[test]
4354 fn metamorphic_mixed_effect_batches_match_sequential_oracle() {
4355 let cases = vec![
4356 synthetic_tool_case(
4357 0,
4358 "read",
4359 Some(ToolEffects::read()),
4360 SyntheticOutcome::Success,
4361 ),
4362 synthetic_tool_case(
4363 1,
4364 "network",
4365 Some(ToolEffects::network()),
4366 SyntheticOutcome::Success,
4367 ),
4368 synthetic_tool_case(
4369 2,
4370 "write",
4371 Some(ToolEffects::write()),
4372 SyntheticOutcome::Success,
4373 ),
4374 synthetic_tool_case(
4375 3,
4376 "read",
4377 Some(ToolEffects::read()),
4378 SyntheticOutcome::Success,
4379 ),
4380 synthetic_tool_case(
4381 4,
4382 "append",
4383 Some(ToolEffects::append()),
4384 SyntheticOutcome::Error,
4385 ),
4386 synthetic_tool_case(
4387 5,
4388 "network",
4389 Some(ToolEffects::network()),
4390 SyntheticOutcome::Success,
4391 ),
4392 synthetic_tool_case(
4393 6,
4394 "process",
4395 Some(ToolEffects::process()),
4396 SyntheticOutcome::Success,
4397 ),
4398 synthetic_tool_case(
4399 7,
4400 "read",
4401 Some(ToolEffects::read()),
4402 SyntheticOutcome::Error,
4403 ),
4404 synthetic_tool_case(8, "unknown", None, SyntheticOutcome::Success),
4405 synthetic_tool_case(
4406 9,
4407 "network",
4408 Some(ToolEffects::network()),
4409 SyntheticOutcome::Success,
4410 ),
4411 ];
4412
4413 assert_eq!(
4414 batch_ranges(&effect_plan(&cases)),
4415 vec![
4416 (0, 2),
4417 (2, 3),
4418 (3, 4),
4419 (4, 5),
4420 (5, 6),
4421 (6, 7),
4422 (7, 8),
4423 (8, 9),
4424 (9, 10)
4425 ]
4426 );
4427 let evidence = tool_effect_batch_plan_evidence(&effect_plan(&cases), 16);
4428 assert_eq!(evidence.schema, TOOL_EFFECT_BATCH_PLAN_SCHEMA_V1);
4429 assert_eq!(evidence.parallelism_cap, 16);
4430 assert_eq!(evidence.batches.len(), 9);
4431 assert!(evidence.batches.iter().any(|batch| {
4432 batch.barrier_reason == Some("append_barrier") && batch.combined_effects == ["append"]
4433 }));
4434 assert!(
4435 cases
4436 .iter()
4437 .any(|case| matches!(case.outcome, SyntheticOutcome::Error)),
4438 "mixed-effect fixture must include failure cases"
4439 );
4440 assert_barrier_effects_are_singleton_batches(&cases);
4441
4442 let oracle = sequential_oracle(&cases);
4443 assert_eq!(
4444 scheduled_transcript(&cases, BatchArrivalOrder::Reverse),
4445 oracle
4446 );
4447 assert_eq!(
4448 scheduled_transcript(&cases, BatchArrivalOrder::RotateLeft(1)),
4449 oracle
4450 );
4451 }
4452
4453 #[test]
4454 fn metamorphic_high_count_batches_keep_transcript_deterministic() {
4455 let cases = (0..96)
4456 .map(|index| match index % 12 {
4457 0 => synthetic_tool_case(
4458 index,
4459 format!("process-{index}"),
4460 Some(ToolEffects::process()),
4461 SyntheticOutcome::Success,
4462 ),
4463 5 => synthetic_tool_case(
4464 index,
4465 format!("append-{index}"),
4466 Some(ToolEffects::append()),
4467 SyntheticOutcome::Success,
4468 ),
4469 9 => synthetic_tool_case(
4470 index,
4471 format!("unknown-{index}"),
4472 None,
4473 SyntheticOutcome::Error,
4474 ),
4475 3 | 7 => synthetic_tool_case(
4476 index,
4477 format!("network-{index}"),
4478 Some(ToolEffects::network()),
4479 SyntheticOutcome::Success,
4480 ),
4481 _ => synthetic_tool_case(
4482 index,
4483 format!("read-{index}"),
4484 Some(ToolEffects::read()),
4485 SyntheticOutcome::Success,
4486 ),
4487 })
4488 .collect::<Vec<_>>();
4489
4490 assert_barrier_effects_are_singleton_batches(&cases);
4491 let oracle = sequential_oracle(&cases);
4492 assert_eq!(
4493 scheduled_transcript(&cases, BatchArrivalOrder::Forward),
4494 oracle
4495 );
4496 assert_eq!(
4497 scheduled_transcript(&cases, BatchArrivalOrder::Reverse),
4498 oracle
4499 );
4500 assert_eq!(
4501 scheduled_transcript(&cases, BatchArrivalOrder::RotateLeft(3)),
4502 oracle
4503 );
4504 }
4505}
4506
4507#[cfg(test)]
4508mod extensions_integration_tests {
4509 use super::*;
4510
4511 use crate::session::Session;
4512 use asupersync::runtime::RuntimeBuilder;
4513 use async_trait::async_trait;
4514 use futures::Stream;
4515 use serde_json::json;
4516 use std::path::Path;
4517 use std::pin::Pin;
4518 use std::sync::atomic::AtomicUsize;
4519 use std::time::Duration;
4520
4521 #[derive(Debug)]
4522 struct NoopProvider;
4523
4524 #[async_trait]
4525 #[allow(clippy::unnecessary_literal_bound)]
4526 impl Provider for NoopProvider {
4527 fn name(&self) -> &str {
4528 "test-provider"
4529 }
4530
4531 fn api(&self) -> &str {
4532 "test-api"
4533 }
4534
4535 fn model_id(&self) -> &str {
4536 "test-model"
4537 }
4538
4539 async fn stream(
4540 &self,
4541 _context: &Context<'_>,
4542 _options: &StreamOptions,
4543 ) -> crate::error::Result<
4544 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
4545 > {
4546 Ok(Box::pin(futures::stream::empty()))
4547 }
4548 }
4549
4550 #[derive(Debug)]
4551 struct IdleCommandProvider;
4552
4553 #[async_trait]
4554 #[allow(clippy::unnecessary_literal_bound)]
4555 impl Provider for IdleCommandProvider {
4556 fn name(&self) -> &str {
4557 "test-provider"
4558 }
4559
4560 fn api(&self) -> &str {
4561 "test-api"
4562 }
4563
4564 fn model_id(&self) -> &str {
4565 "test-model"
4566 }
4567
4568 async fn stream(
4569 &self,
4570 _context: &Context<'_>,
4571 _options: &StreamOptions,
4572 ) -> crate::error::Result<
4573 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
4574 > {
4575 let partial = AssistantMessage {
4576 content: Vec::new(),
4577 api: self.api().to_string(),
4578 provider: self.name().to_string(),
4579 model: self.model_id().to_string(),
4580 usage: Usage::default(),
4581 stop_reason: StopReason::Stop,
4582 error_message: None,
4583 timestamp: 0,
4584 };
4585 let done = AssistantMessage {
4586 content: vec![ContentBlock::Text(TextContent::new(
4587 "resumed-response-0".to_string(),
4588 ))],
4589 api: self.api().to_string(),
4590 provider: self.name().to_string(),
4591 model: self.model_id().to_string(),
4592 usage: Usage::default(),
4593 stop_reason: StopReason::Stop,
4594 error_message: None,
4595 timestamp: 0,
4596 };
4597 Ok(Box::pin(futures::stream::iter(vec![
4598 Ok(StreamEvent::Start { partial }),
4599 Ok(StreamEvent::Done {
4600 reason: StopReason::Stop,
4601 message: done,
4602 }),
4603 ])))
4604 }
4605 }
4606
4607 #[derive(Debug)]
4608 struct CountingTool {
4609 calls: Arc<AtomicUsize>,
4610 }
4611
4612 #[async_trait]
4613 #[allow(clippy::unnecessary_literal_bound)]
4614 impl Tool for CountingTool {
4615 fn name(&self) -> &str {
4616 "count_tool"
4617 }
4618
4619 fn label(&self) -> &str {
4620 "count_tool"
4621 }
4622
4623 fn description(&self) -> &str {
4624 "counting tool"
4625 }
4626
4627 fn parameters(&self) -> serde_json::Value {
4628 json!({ "type": "object" })
4629 }
4630
4631 async fn execute(
4632 &self,
4633 _tool_call_id: &str,
4634 _input: serde_json::Value,
4635 _on_update: Option<Box<dyn Fn(ToolUpdate) + Send + Sync>>,
4636 ) -> Result<ToolOutput> {
4637 self.calls.fetch_add(1, Ordering::SeqCst);
4638 Ok(ToolOutput {
4639 content: vec![ContentBlock::Text(TextContent::new("ok"))],
4640 details: None,
4641 is_error: false,
4642 })
4643 }
4644 }
4645
4646 #[derive(Debug)]
4647 struct ToolUseProvider {
4648 stream_calls: AtomicUsize,
4649 }
4650
4651 impl ToolUseProvider {
4652 const fn new() -> Self {
4653 Self {
4654 stream_calls: AtomicUsize::new(0),
4655 }
4656 }
4657
4658 fn assistant_message(
4659 &self,
4660 stop_reason: StopReason,
4661 content: Vec<ContentBlock>,
4662 ) -> AssistantMessage {
4663 AssistantMessage {
4664 content,
4665 api: self.api().to_string(),
4666 provider: self.name().to_string(),
4667 model: self.model_id().to_string(),
4668 usage: Usage::default(),
4669 stop_reason,
4670 error_message: None,
4671 timestamp: 0,
4672 }
4673 }
4674 }
4675
4676 #[async_trait]
4677 #[allow(clippy::unnecessary_literal_bound)]
4678 impl Provider for ToolUseProvider {
4679 fn name(&self) -> &str {
4680 "test-provider"
4681 }
4682
4683 fn api(&self) -> &str {
4684 "test-api"
4685 }
4686
4687 fn model_id(&self) -> &str {
4688 "test-model"
4689 }
4690
4691 async fn stream(
4692 &self,
4693 _context: &Context<'_>,
4694 _options: &StreamOptions,
4695 ) -> crate::error::Result<
4696 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
4697 > {
4698 let call_index = self.stream_calls.fetch_add(1, Ordering::SeqCst);
4699
4700 let partial = self.assistant_message(StopReason::Stop, Vec::new());
4701
4702 let (reason, message) = if call_index == 0 {
4703 let tool_calls = vec![
4704 ToolCall {
4705 id: "call-1".to_string(),
4706 name: "count_tool".to_string(),
4707 arguments: json!({}),
4708 thought_signature: None,
4709 },
4710 ToolCall {
4711 id: "call-2".to_string(),
4712 name: "count_tool".to_string(),
4713 arguments: json!({}),
4714 thought_signature: None,
4715 },
4716 ];
4717
4718 (
4719 StopReason::ToolUse,
4720 self.assistant_message(
4721 StopReason::ToolUse,
4722 tool_calls
4723 .into_iter()
4724 .map(ContentBlock::ToolCall)
4725 .collect::<Vec<_>>(),
4726 ),
4727 )
4728 } else {
4729 (
4730 StopReason::Stop,
4731 self.assistant_message(
4732 StopReason::Stop,
4733 vec![ContentBlock::Text(TextContent::new("done"))],
4734 ),
4735 )
4736 };
4737
4738 let events = vec![
4739 Ok(StreamEvent::Start { partial }),
4740 Ok(StreamEvent::Done { reason, message }),
4741 ];
4742 Ok(Box::pin(futures::stream::iter(events)))
4743 }
4744 }
4745
4746 #[test]
4747 fn agent_session_enable_extensions_registers_extension_tools() {
4748 let runtime = RuntimeBuilder::current_thread()
4749 .build()
4750 .expect("runtime build");
4751
4752 runtime.block_on(async {
4753 let temp_dir = tempfile::tempdir().expect("tempdir");
4754 let entry_path = temp_dir.path().join("ext.mjs");
4755 std::fs::write(
4756 &entry_path,
4757 r#"
4758 export default function init(pi) {
4759 pi.registerTool({
4760 name: "hello_tool",
4761 label: "hello_tool",
4762 description: "test tool",
4763 parameters: { type: "object", properties: { name: { type: "string" } } },
4764 execute: async (_callId, input, _onUpdate, _abort, ctx) => {
4765 const who = input && input.name ? String(input.name) : "world";
4766 const cwd = ctx && ctx.cwd ? String(ctx.cwd) : "";
4767 return {
4768 content: [{ type: "text", text: `hello ${who}` }],
4769 details: { from: "extension", cwd: cwd },
4770 isError: false
4771 };
4772 }
4773 });
4774 }
4775 "#,
4776 )
4777 .expect("write extension entry");
4778
4779 let provider = Arc::new(NoopProvider);
4780 let tools = ToolRegistry::new(&[], Path::new("."), None);
4781 let agent = Agent::new(provider, tools, AgentConfig::default());
4782 let session = Arc::new(Mutex::new(Session::in_memory()));
4783 let mut agent_session =
4784 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
4785
4786 agent_session
4787 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
4788 .await
4789 .expect("enable extensions");
4790
4791 let tool = agent_session
4792 .agent
4793 .tools
4794 .get("hello_tool")
4795 .expect("hello_tool registered");
4796
4797 let output = tool
4798 .execute("call-1", json!({ "name": "pi" }), None)
4799 .await
4800 .expect("execute tool");
4801
4802 assert!(!output.is_error);
4803 assert!(
4804 matches!(output.content.as_slice(), [ContentBlock::Text(_)]),
4805 "Expected single text content block, got {:?}",
4806 output.content
4807 );
4808 let [ContentBlock::Text(text)] = output.content.as_slice() else {
4809 return;
4810 };
4811 assert_eq!(text.text, "hello pi");
4812
4813 let details = output.details.expect("details present");
4814 assert_eq!(
4815 details.get("from").and_then(serde_json::Value::as_str),
4816 Some("extension")
4817 );
4818 });
4819 }
4820
4821 #[test]
4822 fn agent_session_enable_extensions_with_no_entries_clears_and_is_noop() {
4823 let runtime = RuntimeBuilder::current_thread()
4824 .build()
4825 .expect("runtime build");
4826
4827 runtime.block_on(async {
4828 let temp_dir = tempfile::tempdir().expect("tempdir");
4829 let provider = Arc::new(NoopProvider);
4830 let tools = ToolRegistry::new(&[], Path::new("."), None);
4831 let agent = Agent::new(provider, tools, AgentConfig::default());
4832 let session = Arc::new(Mutex::new(Session::in_memory()));
4833 let mut agent_session =
4834 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
4835
4836 let dummy_manager = ExtensionManager::new();
4838 agent_session.extensions = Some(crate::extensions::ExtensionRegion::new(dummy_manager.clone()));
4839 agent_session.agent.extensions = Some(dummy_manager.clone());
4840 agent_session.extension_queue_modes = Some(Arc::new(std::sync::Mutex::new(ExtensionQueueModeState::new(
4841 QueueMode::OneAtATime,
4842 QueueMode::OneAtATime,
4843 ))));
4844 agent_session.extension_injected_queue = Some(Arc::new(std::sync::Mutex::new(ExtensionInjectedQueue::default())));
4845
4846 agent_session
4847 .enable_extensions(&[], temp_dir.path(), None, &[])
4848 .await
4849 .expect("empty extension list should be a no-op");
4850
4851 assert!(
4852 agent_session.extensions.is_none(),
4853 "no extension region should be created (and existing should be cleared) for an empty extension list"
4854 );
4855 assert!(
4856 agent_session.agent.extensions.is_none(),
4857 "agent should not report extensions active when nothing was requested"
4858 );
4859 assert!(
4860 agent_session.extension_queue_modes.is_none(),
4861 "empty extension list should clear queue mode mirrors"
4862 );
4863 assert!(
4864 agent_session.extension_injected_queue.is_none(),
4865 "empty extension list should clear injected extension queues"
4866 );
4867 });
4868 }
4869
4870 #[test]
4871 fn agent_session_enable_extensions_rejects_mixed_js_and_native_entries() {
4872 let runtime = RuntimeBuilder::current_thread()
4873 .build()
4874 .expect("runtime build");
4875
4876 runtime.block_on(async {
4877 let temp_dir = tempfile::tempdir().expect("tempdir");
4878 let js_entry = temp_dir.path().join("ext.mjs");
4879 let native_entry = temp_dir.path().join("ext.native.json");
4880 std::fs::write(
4881 &js_entry,
4882 r"
4883 export default function init(_pi) {}
4884 ",
4885 )
4886 .expect("write js extension entry");
4887 std::fs::write(&native_entry, "{}").expect("write native extension descriptor");
4888
4889 let provider = Arc::new(NoopProvider);
4890 let tools = ToolRegistry::new(&[], Path::new("."), None);
4891 let agent = Agent::new(provider, tools, AgentConfig::default());
4892 let session = Arc::new(Mutex::new(Session::in_memory()));
4893 let mut agent_session =
4894 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
4895
4896 let err = agent_session
4897 .enable_extensions(&[], temp_dir.path(), None, &[js_entry, native_entry])
4898 .await
4899 .expect_err("mixed extension runtimes should be rejected");
4900 let msg = err.to_string();
4901 assert!(
4902 msg.contains("Mixed extension runtimes are not supported"),
4903 "unexpected mixed-runtime error message: {msg}"
4904 );
4905 });
4906 }
4907
4908 #[test]
4909 fn extension_send_message_persists_custom_message_entry_when_idle() {
4910 let runtime = RuntimeBuilder::current_thread()
4911 .build()
4912 .expect("runtime build");
4913
4914 runtime.block_on(async {
4915 let temp_dir = tempfile::tempdir().expect("tempdir");
4916 let entry_path = temp_dir.path().join("ext.mjs");
4917 std::fs::write(
4918 &entry_path,
4919 r#"
4920 export default function init(pi) {
4921 pi.registerTool({
4922 name: "emit_message",
4923 label: "emit_message",
4924 description: "emit a custom message",
4925 parameters: { type: "object" },
4926 execute: async () => {
4927 pi.sendMessage({
4928 customType: "note",
4929 content: "hello",
4930 display: true,
4931 details: { from: "test" }
4932 }, {});
4933 return { content: [{ type: "text", text: "ok" }], isError: false };
4934 }
4935 });
4936 }
4937 "#,
4938 )
4939 .expect("write extension entry");
4940
4941 let provider = Arc::new(NoopProvider);
4942 let tools = ToolRegistry::new(&[], Path::new("."), None);
4943 let agent = Agent::new(provider, tools, AgentConfig::default());
4944 let session = Arc::new(Mutex::new(Session::in_memory()));
4945 let mut agent_session = AgentSession::new(
4946 agent,
4947 Arc::clone(&session),
4948 false,
4949 ResolvedCompactionSettings::default(),
4950 );
4951
4952 agent_session
4953 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
4954 .await
4955 .expect("enable extensions");
4956
4957 let tool = agent_session
4958 .agent
4959 .tools
4960 .get("emit_message")
4961 .expect("emit_message registered");
4962
4963 let _ = tool
4964 .execute("call-1", json!({}), None)
4965 .await
4966 .expect("execute tool");
4967
4968 let cx = crate::agent_cx::AgentCx::for_request();
4969 let session_guard = session.lock(cx.cx()).await.expect("lock session");
4970 let messages = session_guard.to_messages_for_current_path();
4971
4972 assert!(
4973 messages.iter().any(|msg| {
4974 matches!(
4975 msg,
4976 Message::Custom(CustomMessage { custom_type, content, display, details, .. })
4977 if custom_type == "note"
4978 && content == "hello"
4979 && *display
4980 && details
4981 .as_ref()
4982 .and_then(|v| v.get("from").and_then(Value::as_str))
4983 .is_some_and(|from| from.eq("test"))
4984 )
4985 }),
4986 "expected custom message to be persisted, got {messages:?}"
4987 );
4988 });
4989 }
4990
4991 #[test]
4992 fn extension_send_message_persists_custom_message_entry_when_idle_after_await() {
4993 let runtime = RuntimeBuilder::current_thread()
4994 .build()
4995 .expect("runtime build");
4996
4997 runtime.block_on(async {
4998 let temp_dir = tempfile::tempdir().expect("tempdir");
4999 let entry_path = temp_dir.path().join("ext.mjs");
5000 std::fs::write(
5001 &entry_path,
5002 r#"
5003 export default function init(pi) {
5004 pi.registerTool({
5005 name: "emit_message",
5006 label: "emit_message",
5007 description: "emit a custom message",
5008 parameters: { type: "object" },
5009 execute: async () => {
5010 await Promise.resolve();
5011 pi.sendMessage({
5012 customType: "note",
5013 content: "hello-after-await",
5014 display: true,
5015 details: { from: "test" }
5016 }, {});
5017 return { content: [{ type: "text", text: "ok" }], isError: false };
5018 }
5019 });
5020 }
5021 "#,
5022 )
5023 .expect("write extension entry");
5024
5025 let provider = Arc::new(NoopProvider);
5026 let tools = ToolRegistry::new(&[], Path::new("."), None);
5027 let agent = Agent::new(provider, tools, AgentConfig::default());
5028 let session = Arc::new(Mutex::new(Session::in_memory()));
5029 let mut agent_session = AgentSession::new(
5030 agent,
5031 Arc::clone(&session),
5032 false,
5033 ResolvedCompactionSettings::default(),
5034 );
5035
5036 agent_session
5037 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5038 .await
5039 .expect("enable extensions");
5040
5041 let tool = agent_session
5042 .agent
5043 .tools
5044 .get("emit_message")
5045 .expect("emit_message registered");
5046
5047 let _ = tool
5048 .execute("call-1", json!({}), None)
5049 .await
5050 .expect("execute tool");
5051
5052 let cx = crate::agent_cx::AgentCx::for_request();
5053 let session_guard = session.lock(cx.cx()).await.expect("lock session");
5054 let messages = session_guard.to_messages_for_current_path();
5055
5056 assert!(
5057 messages.iter().any(|msg| {
5058 matches!(
5059 msg,
5060 Message::Custom(CustomMessage { custom_type, content, display, details, .. })
5061 if custom_type == "note"
5062 && content == "hello-after-await"
5063 && *display
5064 && details
5065 .as_ref()
5066 .and_then(|v| v.get("from").and_then(Value::as_str))
5067 .is_some_and(|from| from.eq("test"))
5068 )
5069 }),
5070 "expected custom message to be persisted, got {messages:?}"
5071 );
5072 });
5073 }
5074
5075 #[test]
5076 fn agent_host_actions_send_message_inherits_cancelled_context_when_locked() {
5077 let runtime = RuntimeBuilder::current_thread()
5078 .build()
5079 .expect("runtime build");
5080
5081 runtime.block_on(async {
5082 let session = Arc::new(Mutex::new(Session::in_memory()));
5083 let actions = AgentSessionHostActions {
5084 session: Arc::clone(&session),
5085 injected: Arc::new(StdMutex::new(ExtensionInjectedQueue::default())),
5086 is_streaming: Arc::new(AtomicBool::new(false)),
5087 is_turn_active: Arc::new(AtomicBool::new(false)),
5088 pending_idle_actions: Arc::new(StdMutex::new(VecDeque::new())),
5089 ai_completion: Arc::new(StdMutex::new(ExtensionAiCompletionHostState {
5090 provider: Arc::new(NoopProvider),
5091 stream_options: StreamOptions::default(),
5092 models: Vec::new(),
5093 })),
5094 };
5095
5096 let hold_cx = crate::agent_cx::AgentCx::for_request();
5097 let held_guard = session.lock(hold_cx.cx()).await.expect("lock session");
5098
5099 let ambient_cx = asupersync::Cx::for_testing();
5100 ambient_cx.set_cancel_requested(true);
5101 let _current = asupersync::Cx::set_current(Some(ambient_cx));
5102 let inner = asupersync::time::timeout(
5103 asupersync::time::wall_now(),
5104 Duration::from_millis(100),
5105 actions.send_message(ExtensionSendMessage {
5106 extension_id: Some("ext".to_string()),
5107 custom_type: "note".to_string(),
5108 content: "blocked".to_string(),
5109 display: false,
5110 details: None,
5111 deliver_as: Some(ExtensionDeliverAs::NextTurn),
5112 trigger_turn: false,
5113 }),
5114 )
5115 .await;
5116 let outcome = inner.expect("cancelled helper should finish before timeout");
5117 let err = outcome.expect_err("session append should fail under inherited cancellation");
5118 assert!(
5119 err.to_string().contains("mutex lock cancelled"),
5120 "unexpected error: {err}"
5121 );
5122
5123 drop(held_guard);
5124
5125 let cx = crate::agent_cx::AgentCx::for_request();
5126 let guard = session.lock(cx.cx()).await.expect("lock session");
5127 assert!(
5128 guard.to_messages_for_current_path().is_empty(),
5129 "cancelled send_message should not append a message"
5130 );
5131 });
5132 }
5133
5134 #[derive(Debug, Default)]
5135 struct PiAiCapturedProviderContext {
5136 system_prompt: Option<String>,
5137 messages: Vec<Message>,
5138 }
5139
5140 #[derive(Debug)]
5141 struct PiAiCaptureProvider {
5142 calls: Arc<StdMutex<Vec<PiAiCapturedProviderContext>>>,
5143 }
5144
5145 #[async_trait]
5146 impl Provider for PiAiCaptureProvider {
5147 fn name(&self) -> &'static str {
5148 "capturing-provider"
5149 }
5150
5151 fn api(&self) -> &'static str {
5152 "test-api"
5153 }
5154
5155 fn model_id(&self) -> &'static str {
5156 "capture-model"
5157 }
5158
5159 async fn stream(
5160 &self,
5161 context: &Context<'_>,
5162 _options: &StreamOptions,
5163 ) -> crate::error::Result<
5164 std::pin::Pin<
5165 Box<dyn futures::Stream<Item = crate::error::Result<StreamEvent>> + Send>,
5166 >,
5167 > {
5168 self.calls
5169 .lock()
5170 .unwrap_or_else(std::sync::PoisonError::into_inner)
5171 .push(PiAiCapturedProviderContext {
5172 system_prompt: context.system_prompt.as_ref().map(ToString::to_string),
5173 messages: context.messages.iter().cloned().collect(),
5174 });
5175 let final_message = AssistantMessage {
5176 content: vec![ContentBlock::Text(TextContent::new("captured"))],
5177 api: "test-api".to_string(),
5178 provider: "capturing-provider".to_string(),
5179 model: "capture-model".to_string(),
5180 usage: Usage::default(),
5181 stop_reason: StopReason::Stop,
5182 error_message: None,
5183 timestamp: 0,
5184 };
5185 Ok(Box::pin(futures::stream::iter(vec![Ok(
5186 StreamEvent::Done {
5187 reason: StopReason::Stop,
5188 message: final_message,
5189 },
5190 )])))
5191 }
5192 }
5193
5194 #[test]
5195 fn agent_host_actions_complete_ai_streams_configured_provider() {
5196 let runtime = RuntimeBuilder::current_thread()
5197 .build()
5198 .expect("runtime build");
5199
5200 runtime.block_on(async {
5201 let session = Arc::new(Mutex::new(Session::in_memory()));
5202 let calls = Arc::new(StdMutex::new(Vec::new()));
5203 let provider = Arc::new(PiAiCaptureProvider {
5204 calls: Arc::clone(&calls),
5205 });
5206 let actions = AgentSessionHostActions {
5207 session,
5208 injected: Arc::new(StdMutex::new(ExtensionInjectedQueue::default())),
5209 is_streaming: Arc::new(AtomicBool::new(false)),
5210 is_turn_active: Arc::new(AtomicBool::new(false)),
5211 pending_idle_actions: Arc::new(StdMutex::new(VecDeque::new())),
5212 ai_completion: Arc::new(StdMutex::new(ExtensionAiCompletionHostState {
5213 provider,
5214 stream_options: StreamOptions::default(),
5215 models: vec![json!({
5216 "id": "capture-model",
5217 "provider": "capturing-provider",
5218 "api": "test-api",
5219 })],
5220 })),
5221 };
5222
5223 let result = actions
5224 .complete_ai(ExtensionAiCompletionRequest {
5225 model: json!({ "id": "capture-model" }),
5226 context: json!({
5227 "systemPrompt": "answer tersely",
5228 "messages": [
5229 { "role": "user", "content": "ping" }
5230 ]
5231 }),
5232 options: json!({ "maxTokens": 16 }),
5233 simple: false,
5234 })
5235 .await
5236 .expect("complete through provider");
5237
5238 assert_eq!(result["text"], json!("captured"));
5239 assert_eq!(result["provider"], json!("capturing-provider"));
5240 assert_eq!(result["api"], json!("test-api"));
5241
5242 let (captured_len, captured_system_prompt, captured_messages) = {
5243 let captured = match calls.lock() {
5244 Ok(guard) => guard,
5245 Err(poisoned) => poisoned.into_inner(),
5246 };
5247 (
5248 captured.len(),
5249 captured.first().and_then(|call| call.system_prompt.clone()),
5250 captured
5251 .first()
5252 .map(|call| call.messages.clone())
5253 .unwrap_or_default(),
5254 )
5255 };
5256 assert_eq!(captured_len, 1);
5257 assert_eq!(captured_system_prompt.as_deref(), Some("answer tersely"));
5258 assert_eq!(captured_messages.len(), 1);
5259 assert!(
5260 matches!(
5261 captured_messages.first(),
5262 Some(Message::User(UserMessage { content: UserContent::Text(text), .. }))
5263 if text == "ping"
5264 ),
5265 "expected user message context, got {captured_messages:?}"
5266 );
5267
5268 let models = actions.list_ai_models().await.expect("list models");
5269 assert_eq!(models[0]["id"], json!("capture-model"));
5270 });
5271 }
5272
5273 #[test]
5274 fn extension_command_send_message_trigger_turn_runs_agent_turn_when_idle() {
5275 let runtime = RuntimeBuilder::current_thread()
5276 .build()
5277 .expect("runtime build");
5278
5279 runtime.block_on(async {
5280 let temp_dir = tempfile::tempdir().expect("tempdir");
5281 let entry_path = temp_dir.path().join("ext.mjs");
5282 std::fs::write(
5283 &entry_path,
5284 r#"
5285 export default function init(pi) {
5286 pi.registerCommand("emit-now", {
5287 description: "emit a custom message and trigger a turn",
5288 handler: async () => {
5289 await pi.events("sendMessage", {
5290 message: {
5291 customType: "note",
5292 content: "turn-now",
5293 display: true
5294 },
5295 options: {
5296 deliverAs: "steer",
5297 triggerTurn: true
5298 }
5299 });
5300 return "queued";
5301 }
5302 });
5303 }
5304 "#,
5305 )
5306 .expect("write extension entry");
5307
5308 let provider = Arc::new(IdleCommandProvider);
5309 let tools = ToolRegistry::new(&[], Path::new("."), None);
5310 let agent = Agent::new(provider, tools, AgentConfig::default());
5311 let session = Arc::new(Mutex::new(Session::in_memory()));
5312 let mut agent_session = AgentSession::new(
5313 agent,
5314 Arc::clone(&session),
5315 false,
5316 ResolvedCompactionSettings::default(),
5317 );
5318
5319 agent_session
5320 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5321 .await
5322 .expect("enable extensions");
5323
5324 let value = agent_session
5325 .execute_extension_command("emit-now", "", 5_000, |_| {})
5326 .await
5327 .expect("execute extension command");
5328 assert_eq!(value.as_str(), Some("queued"));
5329
5330 let cx = crate::agent_cx::AgentCx::for_request();
5331 let session_guard = session.lock(cx.cx()).await.expect("lock session");
5332 let messages = session_guard.to_messages_for_current_path();
5333
5334 assert!(
5335 messages.iter().any(|msg| {
5336 matches!(
5337 msg,
5338 Message::Custom(CustomMessage { custom_type, content, .. })
5339 if custom_type == "note" && content == "turn-now"
5340 )
5341 }),
5342 "expected custom message prompt in session, got {messages:?}"
5343 );
5344 assert!(
5345 messages.iter().any(|msg| {
5346 matches!(
5347 msg,
5348 Message::Assistant(assistant)
5349 if assistant.content.iter().any(|block| matches!(
5350 block,
5351 ContentBlock::Text(TextContent { text, .. })
5352 if text.as_str().eq("resumed-response-0")
5353 ))
5354 )
5355 }),
5356 "expected assistant response after triggered turn, got {messages:?}"
5357 );
5358 });
5359 }
5360
5361 #[test]
5362 fn agent_extension_session_get_state_reports_agent_runtime_state() {
5363 let runtime = RuntimeBuilder::current_thread()
5364 .build()
5365 .expect("runtime build");
5366
5367 runtime.block_on(async {
5368 let mut session = Session::in_memory();
5369 session.set_model_header(
5370 Some("test-provider".to_string()),
5371 Some("test-model".to_string()),
5372 Some("high".to_string()),
5373 );
5374 session.append_message(crate::session::SessionMessage::User {
5375 content: UserContent::Text("hello".to_string()),
5376 timestamp: Some(1),
5377 });
5378 let session = Arc::new(Mutex::new(session));
5379
5380 let extension_session = AgentExtensionSession {
5381 handle: SessionHandle(Arc::clone(&session)),
5382 is_streaming: Arc::new(AtomicBool::new(true)),
5383 is_compacting: Arc::new(AtomicBool::new(true)),
5384 queue_modes: Arc::new(StdMutex::new(ExtensionQueueModeState::new(
5385 QueueMode::All,
5386 QueueMode::OneAtATime,
5387 ))),
5388 auto_compaction_enabled: true,
5389 };
5390
5391 let state = <AgentExtensionSession as crate::extensions::ExtensionSession>::get_state(
5392 &extension_session,
5393 )
5394 .await;
5395
5396 assert_eq!(state["model"]["provider"], "test-provider");
5397 assert_eq!(state["model"]["id"], "test-model");
5398 assert_eq!(state["thinkingLevel"], "high");
5399 assert_eq!(state["isStreaming"], true);
5400 assert_eq!(state["isCompacting"], true);
5401 assert_eq!(state["steeringMode"], "all");
5402 assert_eq!(state["followUpMode"], "one-at-a-time");
5403 assert_eq!(state["autoCompactionEnabled"], true);
5404 assert_eq!(state["messageCount"], 1);
5405 });
5406 }
5407
5408 #[test]
5409 fn agent_extension_session_get_state_uses_branch_local_model_and_thinking() {
5410 let runtime = RuntimeBuilder::current_thread()
5411 .build()
5412 .expect("runtime build");
5413
5414 runtime.block_on(async {
5415 let mut session = Session::in_memory();
5416 let root_id = session.append_message(crate::session::SessionMessage::User {
5417 content: UserContent::Text("root".to_string()),
5418 timestamp: Some(1),
5419 });
5420 session.append_model_change("openai".to_string(), "gpt-4o".to_string());
5421 let branch_a_thinking = session.append_thinking_level_change("low".to_string());
5422 session.set_model_header(
5423 Some("openai".to_string()),
5424 Some("gpt-4o".to_string()),
5425 Some("low".to_string()),
5426 );
5427
5428 assert!(session.create_branch_from(&root_id));
5429 session.append_model_change("anthropic".to_string(), "claude-sonnet-4-5".to_string());
5430 session.append_thinking_level_change("high".to_string());
5431 session.set_model_header(
5432 Some("anthropic".to_string()),
5433 Some("claude-sonnet-4-5".to_string()),
5434 Some("high".to_string()),
5435 );
5436
5437 assert!(session.navigate_to(&branch_a_thinking));
5438 let session = Arc::new(Mutex::new(session));
5439
5440 let extension_session = AgentExtensionSession {
5441 handle: SessionHandle(Arc::clone(&session)),
5442 is_streaming: Arc::new(AtomicBool::new(false)),
5443 is_compacting: Arc::new(AtomicBool::new(false)),
5444 queue_modes: Arc::new(StdMutex::new(ExtensionQueueModeState::new(
5445 QueueMode::OneAtATime,
5446 QueueMode::OneAtATime,
5447 ))),
5448 auto_compaction_enabled: false,
5449 };
5450
5451 let state = <AgentExtensionSession as crate::extensions::ExtensionSession>::get_state(
5452 &extension_session,
5453 )
5454 .await;
5455
5456 assert_eq!(state["model"]["provider"], "openai");
5457 assert_eq!(state["model"]["id"], "gpt-4o");
5458 assert_eq!(state["thinkingLevel"], "low");
5459 });
5460 }
5461
5462 #[test]
5463 fn agent_session_set_queue_modes_updates_extension_delivery_state() {
5464 let provider = Arc::new(NoopProvider);
5465 let tools = ToolRegistry::new(&[], Path::new("."), None);
5466 let agent = Agent::new(provider, tools, AgentConfig::default());
5467 let session = Arc::new(Mutex::new(Session::in_memory()));
5468 let mut agent_session =
5469 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
5470
5471 let queue_modes = Arc::new(StdMutex::new(ExtensionQueueModeState::new(
5472 QueueMode::OneAtATime,
5473 QueueMode::OneAtATime,
5474 )));
5475 let injected_queue = Arc::new(StdMutex::new(ExtensionInjectedQueue::new(
5476 QueueMode::OneAtATime,
5477 QueueMode::OneAtATime,
5478 )));
5479 agent_session.extension_queue_modes = Some(Arc::clone(&queue_modes));
5480 agent_session.extension_injected_queue = Some(Arc::clone(&injected_queue));
5481
5482 agent_session.set_queue_modes(QueueMode::All, QueueMode::All);
5483
5484 assert_eq!(
5485 agent_session.agent.queue_modes(),
5486 (QueueMode::All, QueueMode::All)
5487 );
5488 let mirrored = queue_modes.lock().expect("lock queue mode mirror");
5489 assert_eq!(mirrored.steering_mode, QueueMode::All);
5490 assert_eq!(mirrored.follow_up_mode, QueueMode::All);
5491 drop(mirrored);
5492
5493 let queued_follow_up_len = {
5494 let mut queue = injected_queue.lock().expect("lock injected queue");
5495 queue.push_follow_up(Message::User(UserMessage {
5496 content: UserContent::Text("first".to_string()),
5497 timestamp: 0,
5498 }));
5499 queue.push_follow_up(Message::User(UserMessage {
5500 content: UserContent::Text("second".to_string()),
5501 timestamp: 0,
5502 }));
5503 queue.pop_follow_up().len()
5504 };
5505 assert_eq!(
5506 queued_follow_up_len, 2,
5507 "updated queue modes should apply to extension-injected follow-ups"
5508 );
5509 }
5510
5511 #[test]
5512 fn extension_command_send_user_message_runs_agent_turn_when_idle() {
5513 let runtime = RuntimeBuilder::current_thread()
5514 .build()
5515 .expect("runtime build");
5516
5517 runtime.block_on(async {
5518 let temp_dir = tempfile::tempdir().expect("tempdir");
5519 let entry_path = temp_dir.path().join("ext.mjs");
5520 std::fs::write(
5521 &entry_path,
5522 r#"
5523 export default function init(pi) {
5524 pi.registerCommand("inject-user", {
5525 description: "inject a user message",
5526 handler: async () => {
5527 await pi.events("sendUserMessage", {
5528 text: "Please review the changes"
5529 });
5530 return "queued";
5531 }
5532 });
5533 }
5534 "#,
5535 )
5536 .expect("write extension entry");
5537
5538 let provider = Arc::new(IdleCommandProvider);
5539 let tools = ToolRegistry::new(&[], Path::new("."), None);
5540 let agent = Agent::new(provider, tools, AgentConfig::default());
5541 let session = Arc::new(Mutex::new(Session::in_memory()));
5542 let mut agent_session = AgentSession::new(
5543 agent,
5544 Arc::clone(&session),
5545 false,
5546 ResolvedCompactionSettings::default(),
5547 );
5548
5549 agent_session
5550 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5551 .await
5552 .expect("enable extensions");
5553
5554 let value = agent_session
5555 .execute_extension_command("inject-user", "", 5_000, |_| {})
5556 .await
5557 .expect("execute extension command");
5558 assert_eq!(value.as_str(), Some("queued"));
5559
5560 let cx = crate::agent_cx::AgentCx::for_request();
5561 let session_guard = session.lock(cx.cx()).await.expect("lock session");
5562 let messages = session_guard.to_messages_for_current_path();
5563
5564 assert!(
5565 messages.iter().any(|msg| {
5566 matches!(
5567 msg,
5568 Message::User(UserMessage {
5569 content: UserContent::Text(text),
5570 ..
5571 }) if text == "Please review the changes"
5572 )
5573 }),
5574 "expected injected user message in session, got {messages:?}"
5575 );
5576 assert!(
5577 messages.iter().any(|msg| {
5578 matches!(
5579 msg,
5580 Message::Assistant(assistant)
5581 if assistant.content.iter().any(|block| matches!(
5582 block,
5583 ContentBlock::Text(TextContent { text, .. })
5584 if text.as_str().eq("resumed-response-0")
5585 ))
5586 )
5587 }),
5588 "expected assistant response after injected user turn, got {messages:?}"
5589 );
5590 });
5591 }
5592
5593 #[test]
5594 fn send_user_message_steer_skips_remaining_tools() {
5595 let runtime = RuntimeBuilder::current_thread()
5596 .build()
5597 .expect("runtime build");
5598
5599 runtime.block_on(async {
5600 let temp_dir = tempfile::tempdir().expect("tempdir");
5601 let entry_path = temp_dir.path().join("ext.mjs");
5602 std::fs::write(
5603 &entry_path,
5604 r#"
5605 export default function init(pi) {
5606 let sent = false;
5607 pi.on("tool_call", async (event) => {
5608 if (sent) return {};
5609 if (Object.is(event && event.toolName, "count_tool")) {
5610 sent = true;
5611 await pi.events("sendUserMessage", {
5612 text: "steer-now",
5613 options: { deliverAs: "steer" }
5614 });
5615 }
5616 return {};
5617 });
5618 }
5619 "#,
5620 )
5621 .expect("write extension entry");
5622
5623 let provider = Arc::new(ToolUseProvider::new());
5624 let calls = Arc::new(AtomicUsize::new(0));
5625 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
5626 calls: Arc::clone(&calls),
5627 })]);
5628 let agent = Agent::new(provider, tools, AgentConfig::default());
5629 let session = Arc::new(Mutex::new(Session::in_memory()));
5630 let mut agent_session =
5631 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
5632
5633 agent_session
5634 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5635 .await
5636 .expect("enable extensions");
5637
5638 let _ = agent_session
5639 .run_text("go".to_string(), |_| {})
5640 .await
5641 .expect("run_text");
5642
5643 assert_eq!(calls.load(Ordering::SeqCst), 1);
5645 });
5646 }
5647
5648 #[test]
5649 fn send_user_message_follow_up_does_not_skip_tools() {
5650 let runtime = RuntimeBuilder::current_thread()
5651 .build()
5652 .expect("runtime build");
5653
5654 runtime.block_on(async {
5655 let temp_dir = tempfile::tempdir().expect("tempdir");
5656 let entry_path = temp_dir.path().join("ext.mjs");
5657 std::fs::write(
5658 &entry_path,
5659 r#"
5660 export default function init(pi) {
5661 let sent = false;
5662 pi.on("tool_call", async (event) => {
5663 if (sent) return {};
5664 if (Object.is(event && event.toolName, "count_tool")) {
5665 sent = true;
5666 await pi.events("sendUserMessage", {
5667 text: "follow-up",
5668 options: { deliverAs: "followUp" }
5669 });
5670 }
5671 return {};
5672 });
5673 }
5674 "#,
5675 )
5676 .expect("write extension entry");
5677
5678 let provider = Arc::new(ToolUseProvider::new());
5679 let calls = Arc::new(AtomicUsize::new(0));
5680 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
5681 calls: Arc::clone(&calls),
5682 })]);
5683 let agent = Agent::new(provider, tools, AgentConfig::default());
5684 let session = Arc::new(Mutex::new(Session::in_memory()));
5685 let mut agent_session =
5686 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
5687
5688 agent_session
5689 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5690 .await
5691 .expect("enable extensions");
5692
5693 let _ = agent_session
5694 .run_text("go".to_string(), |_| {})
5695 .await
5696 .expect("run_text");
5697
5698 assert_eq!(calls.load(Ordering::SeqCst), 2);
5699 });
5700 }
5701
5702 fn test_turn_latency() -> SharedTurnLatencyAccumulator {
5703 Arc::new(StdMutex::new(TurnLatencyAccumulator::started()))
5704 }
5705
5706 #[test]
5707 fn latency_breakdown_reports_component_tail_percentiles() {
5708 let breakdown =
5709 TurnLatencyBreakdown::from_component_samples(250, &[10, 30, 20], &[40, 5], &[2], &[]);
5710
5711 assert_eq!(breakdown.schema, TURN_LATENCY_BREAKDOWN_SCHEMA_V1);
5712 assert_eq!(breakdown.provider_streaming.duration_ms, 60);
5713 assert_eq!(breakdown.provider_streaming.samples, 3);
5714 assert_eq!(breakdown.provider_streaming.tail_percentiles.p50_ms, 20);
5715 assert_eq!(breakdown.provider_streaming.tail_percentiles.p95_ms, 30);
5716 assert_eq!(breakdown.provider_streaming.tail_percentiles.p99_ms, 30);
5717 assert_eq!(breakdown.provider_streaming.tail_percentiles.p999_ms, 30);
5718 assert_eq!(breakdown.local_tools.duration_ms, 45);
5719 assert_eq!(breakdown.extension_hostcalls.duration_ms, 2);
5720 assert_eq!(breakdown.persistence.duration_ms, 0);
5721 assert_eq!(breakdown.dominant_component, "provider_streaming");
5722 }
5723
5724 #[test]
5725 fn latency_breakdown_serializes_without_provider_secrets() {
5726 let breakdown =
5727 TurnLatencyBreakdown::from_component_samples(125, &[100], &[20], &[5], &[0]);
5728 let serialized = serde_json::to_string(&breakdown).expect("serialize latency breakdown");
5729
5730 assert!(serialized.contains(TURN_LATENCY_BREAKDOWN_SCHEMA_V1));
5731 assert!(serialized.contains("providerStreaming"));
5732 assert!(serialized.contains("localTools"));
5733 assert!(serialized.contains("extensionHostcalls"));
5734 assert!(serialized.contains("persistence"));
5735 assert!(!serialized.contains("api_key"));
5736 assert!(!serialized.contains("authorization"));
5737 assert!(!serialized.contains("bearer"));
5738 assert!(!serialized.contains("sk-"));
5739 }
5740
5741 #[test]
5742 fn tool_call_hook_can_block_tool_execution() {
5743 let runtime = RuntimeBuilder::current_thread()
5744 .build()
5745 .expect("runtime build");
5746
5747 runtime.block_on(async {
5748 let temp_dir = tempfile::tempdir().expect("tempdir");
5749 let entry_path = temp_dir.path().join("ext.mjs");
5750 std::fs::write(
5751 &entry_path,
5752 r#"
5753 export default function init(pi) {
5754 pi.on("tool_call", async (event) => {
5755 if (Object.is(event && event.toolName, "count_tool")) {
5756 return { block: true, reason: "blocked in test" };
5757 }
5758 return {};
5759 });
5760 }
5761 "#,
5762 )
5763 .expect("write extension entry");
5764
5765 let provider = Arc::new(NoopProvider);
5766 let calls = Arc::new(AtomicUsize::new(0));
5767 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
5768 calls: Arc::clone(&calls),
5769 })]);
5770 let agent = Agent::new(provider, tools, AgentConfig::default());
5771 let session = Arc::new(Mutex::new(Session::in_memory()));
5772 let mut agent_session =
5773 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
5774
5775 agent_session
5776 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5777 .await
5778 .expect("enable extensions");
5779
5780 let tool_call = ToolCall {
5781 id: "call-1".to_string(),
5782 name: "count_tool".to_string(),
5783 arguments: json!({}),
5784 thought_signature: None,
5785 };
5786
5787 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
5788 let (output, is_error) = agent_session
5789 .agent
5790 .execute_tool(tool_call, on_event, test_turn_latency())
5791 .await;
5792
5793 assert!(is_error);
5794 assert!(output.is_error);
5795 assert_eq!(calls.load(Ordering::SeqCst), 0);
5796
5797 assert_eq!(output.details, None);
5798 assert!(
5799 matches!(output.content.as_slice(), [ContentBlock::Text(_)]),
5800 "Expected text output, got {:?}",
5801 output.content
5802 );
5803 if let [ContentBlock::Text(text)] = output.content.as_slice() {
5804 assert_eq!(text.text, "Tool execution blocked: blocked in test");
5805 }
5806 });
5807 }
5808
5809 #[test]
5810 fn tool_call_hook_errors_fail_open() {
5811 let runtime = RuntimeBuilder::current_thread()
5812 .build()
5813 .expect("runtime build");
5814
5815 runtime.block_on(async {
5816 let temp_dir = tempfile::tempdir().expect("tempdir");
5817 let entry_path = temp_dir.path().join("ext.mjs");
5818 std::fs::write(
5819 &entry_path,
5820 r#"
5821 export default function init(pi) {
5822 pi.on("tool_call", async (_event) => {
5823 throw new Error("boom");
5824 });
5825 }
5826 "#,
5827 )
5828 .expect("write extension entry");
5829
5830 let provider = Arc::new(NoopProvider);
5831 let calls = Arc::new(AtomicUsize::new(0));
5832 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
5833 calls: Arc::clone(&calls),
5834 })]);
5835 let agent = Agent::new(provider, tools, AgentConfig::default());
5836 let session = Arc::new(Mutex::new(Session::in_memory()));
5837 let mut agent_session =
5838 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
5839
5840 agent_session
5841 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5842 .await
5843 .expect("enable extensions");
5844
5845 let tool_call = ToolCall {
5846 id: "call-1".to_string(),
5847 name: "count_tool".to_string(),
5848 arguments: json!({}),
5849 thought_signature: None,
5850 };
5851
5852 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
5853 let (output, is_error) = agent_session
5854 .agent
5855 .execute_tool(tool_call, on_event, test_turn_latency())
5856 .await;
5857
5858 assert!(!is_error);
5859 assert!(!output.is_error);
5860 assert_eq!(calls.load(Ordering::SeqCst), 1);
5861 });
5862 }
5863
5864 #[test]
5865 fn tool_call_hook_errors_fail_closed_when_configured() {
5866 let runtime = RuntimeBuilder::current_thread()
5867 .build()
5868 .expect("runtime build");
5869
5870 runtime.block_on(async {
5871 let temp_dir = tempfile::tempdir().expect("tempdir");
5872 let entry_path = temp_dir.path().join("ext.mjs");
5873 std::fs::write(
5874 &entry_path,
5875 r#"
5876 export default function init(pi) {
5877 pi.on("tool_call", async (_event) => {
5878 throw new Error("boom");
5879 });
5880 }
5881 "#,
5882 )
5883 .expect("write extension entry");
5884
5885 let provider = Arc::new(NoopProvider);
5886 let calls = Arc::new(AtomicUsize::new(0));
5887 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
5888 calls: Arc::clone(&calls),
5889 })]);
5890 let agent = Agent::new(
5891 provider,
5892 tools,
5893 AgentConfig {
5894 fail_closed_hooks: true,
5895 ..AgentConfig::default()
5896 },
5897 );
5898 let session = Arc::new(Mutex::new(Session::in_memory()));
5899 let mut agent_session =
5900 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
5901
5902 agent_session
5903 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5904 .await
5905 .expect("enable extensions");
5906
5907 let tool_call = ToolCall {
5908 id: "call-1".to_string(),
5909 name: "count_tool".to_string(),
5910 arguments: json!({}),
5911 thought_signature: None,
5912 };
5913
5914 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
5915 let (output, is_error) = agent_session
5916 .agent
5917 .execute_tool(tool_call, on_event, test_turn_latency())
5918 .await;
5919
5920 assert!(is_error);
5921 assert!(output.is_error);
5922 assert_eq!(calls.load(Ordering::SeqCst), 0);
5923 assert!(
5924 matches!(output.content.as_slice(), [ContentBlock::Text(_)]),
5925 "Expected text output, got {:?}",
5926 output.content
5927 );
5928 let [ContentBlock::Text(text)] = output.content.as_slice() else {
5929 return;
5930 };
5931 assert_eq!(text.text, "Tool execution blocked: extension hook failed");
5932 });
5933 }
5934
5935 #[test]
5936 fn tool_call_hook_absent_allows_tool_execution() {
5937 let runtime = RuntimeBuilder::current_thread()
5938 .build()
5939 .expect("runtime build");
5940
5941 runtime.block_on(async {
5942 let temp_dir = tempfile::tempdir().expect("tempdir");
5943 let entry_path = temp_dir.path().join("ext.mjs");
5944 std::fs::write(
5945 &entry_path,
5946 r"
5947 export default function init(_pi) {}
5948 ",
5949 )
5950 .expect("write extension entry");
5951
5952 let provider = Arc::new(NoopProvider);
5953 let calls = Arc::new(AtomicUsize::new(0));
5954 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
5955 calls: Arc::clone(&calls),
5956 })]);
5957 let agent = Agent::new(provider, tools, AgentConfig::default());
5958 let session = Arc::new(Mutex::new(Session::in_memory()));
5959 let mut agent_session =
5960 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
5961
5962 agent_session
5963 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
5964 .await
5965 .expect("enable extensions");
5966
5967 let tool_call = ToolCall {
5968 id: "call-1".to_string(),
5969 name: "count_tool".to_string(),
5970 arguments: json!({}),
5971 thought_signature: None,
5972 };
5973
5974 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
5975 let (output, is_error) = agent_session
5976 .agent
5977 .execute_tool(tool_call, on_event, test_turn_latency())
5978 .await;
5979
5980 assert!(!is_error);
5981 assert!(!output.is_error);
5982 assert_eq!(calls.load(Ordering::SeqCst), 1);
5983 });
5984 }
5985
5986 #[test]
5987 fn tool_approval_allow_executes_tool() {
5988 let runtime = RuntimeBuilder::current_thread()
5989 .build()
5990 .expect("runtime build");
5991
5992 runtime.block_on(async {
5993 let provider = Arc::new(NoopProvider);
5994 let calls = Arc::new(AtomicUsize::new(0));
5995 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
5996 calls: Arc::clone(&calls),
5997 })]);
5998 let approval_calls = Arc::new(AtomicUsize::new(0));
5999 let approval_counter = Arc::clone(&approval_calls);
6000 let agent = Agent::new(
6001 provider,
6002 tools,
6003 AgentConfig {
6004 tool_approval: Some(Arc::new(move |request| {
6005 assert_eq!(request.tool_call_id, "call-1");
6006 assert_eq!(request.tool_name, "count_tool");
6007 approval_counter.fetch_add(1, Ordering::SeqCst);
6008 Box::pin(async { ToolApprovalDecision::Allow })
6009 })),
6010 ..AgentConfig::default()
6011 },
6012 );
6013 let session = Arc::new(Mutex::new(Session::in_memory()));
6014 let agent_session =
6015 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6016
6017 let tool_call = ToolCall {
6018 id: "call-1".to_string(),
6019 name: "count_tool".to_string(),
6020 arguments: json!({}),
6021 thought_signature: None,
6022 };
6023
6024 let events = Arc::new(std::sync::Mutex::new(Vec::new()));
6025 let events_for_handler = Arc::clone(&events);
6026 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(move |event| {
6027 if let Ok(mut guard) = events_for_handler.lock() {
6028 guard.push(event);
6029 }
6030 });
6031 let (output, is_error) = agent_session
6032 .agent
6033 .execute_tool(tool_call, on_event, test_turn_latency())
6034 .await;
6035
6036 assert!(!is_error);
6037 assert!(!output.is_error);
6038 assert_eq!(approval_calls.load(Ordering::SeqCst), 1);
6039 assert_eq!(calls.load(Ordering::SeqCst), 1);
6040 let saw_approval_update = events.lock().is_ok_and(|guard| {
6041 guard.iter().any(|event| {
6042 matches!(
6043 event,
6044 AgentEvent::ToolExecutionUpdate {
6045 partial_result,
6046 ..
6047 } if partial_result.details.as_ref().is_some_and(|details| {
6048 details["schema"] == TOOL_APPROVAL_STATUS_SCHEMA_V1
6049 && details["status"] == "approved"
6050 })
6051 )
6052 })
6053 });
6054 assert!(saw_approval_update);
6055 });
6056 }
6057
6058 #[test]
6059 fn tool_approval_deny_blocks_tool_execution() {
6060 let runtime = RuntimeBuilder::current_thread()
6061 .build()
6062 .expect("runtime build");
6063
6064 runtime.block_on(async {
6065 let provider = Arc::new(NoopProvider);
6066 let calls = Arc::new(AtomicUsize::new(0));
6067 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
6068 calls: Arc::clone(&calls),
6069 })]);
6070 let agent = Agent::new(
6071 provider,
6072 tools,
6073 AgentConfig {
6074 tool_approval: Some(Arc::new(|request| {
6075 assert_eq!(request.tool_name, "count_tool");
6076 Box::pin(async { ToolApprovalDecision::deny("denied by approval test") })
6077 })),
6078 ..AgentConfig::default()
6079 },
6080 );
6081 let session = Arc::new(Mutex::new(Session::in_memory()));
6082 let agent_session =
6083 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6084
6085 let tool_call = ToolCall {
6086 id: "call-1".to_string(),
6087 name: "count_tool".to_string(),
6088 arguments: json!({}),
6089 thought_signature: None,
6090 };
6091
6092 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
6093 let (output, is_error) = agent_session
6094 .agent
6095 .execute_tool(tool_call, on_event, test_turn_latency())
6096 .await;
6097
6098 assert!(is_error);
6099 assert!(output.is_error);
6100 assert_eq!(calls.load(Ordering::SeqCst), 0);
6101 assert_eq!(
6102 output.details.as_ref().unwrap()["schema"],
6103 TOOL_APPROVAL_DENIED_SCHEMA_V1
6104 );
6105 assert!(
6106 matches!(output.content.as_slice(), [ContentBlock::Text(text)] if text
6107 .text
6108 .contains("denied by approval test"))
6109 );
6110 });
6111 }
6112
6113 #[test]
6114 fn tool_call_hook_returns_empty_allows_tool_execution() {
6115 let runtime = RuntimeBuilder::current_thread()
6116 .build()
6117 .expect("runtime build");
6118
6119 runtime.block_on(async {
6120 let temp_dir = tempfile::tempdir().expect("tempdir");
6121 let entry_path = temp_dir.path().join("ext.mjs");
6122 std::fs::write(
6123 &entry_path,
6124 r#"
6125 export default function init(pi) {
6126 pi.on("tool_call", async (_event) => ({}));
6127 }
6128 "#,
6129 )
6130 .expect("write extension entry");
6131
6132 let provider = Arc::new(NoopProvider);
6133 let calls = Arc::new(AtomicUsize::new(0));
6134 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
6135 calls: Arc::clone(&calls),
6136 })]);
6137 let agent = Agent::new(provider, tools, AgentConfig::default());
6138 let session = Arc::new(Mutex::new(Session::in_memory()));
6139 let mut agent_session =
6140 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6141
6142 agent_session
6143 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
6144 .await
6145 .expect("enable extensions");
6146
6147 let tool_call = ToolCall {
6148 id: "call-1".to_string(),
6149 name: "count_tool".to_string(),
6150 arguments: json!({}),
6151 thought_signature: None,
6152 };
6153
6154 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
6155 let (output, is_error) = agent_session
6156 .agent
6157 .execute_tool(tool_call, on_event, test_turn_latency())
6158 .await;
6159
6160 assert!(!is_error);
6161 assert!(!output.is_error);
6162 assert_eq!(calls.load(Ordering::SeqCst), 1);
6163 });
6164 }
6165
6166 #[test]
6167 fn tool_call_hook_can_block_bash_tool_execution() {
6168 let runtime = RuntimeBuilder::current_thread()
6169 .build()
6170 .expect("runtime build");
6171
6172 runtime.block_on(async {
6173 let temp_dir = tempfile::tempdir().expect("tempdir");
6174 let entry_path = temp_dir.path().join("ext.mjs");
6175 std::fs::write(
6176 &entry_path,
6177 r#"
6178 export default function init(pi) {
6179 pi.on("tool_call", async (event) => {
6180 const name = event && event.toolName ? String(event.toolName) : "";
6181 if (name === "bash") return { block: true, reason: "blocked bash in test" };
6182 return {};
6183 });
6184 }
6185 "#,
6186 )
6187 .expect("write extension entry");
6188
6189 let provider = Arc::new(NoopProvider);
6190 let tools = ToolRegistry::new(&["bash"], temp_dir.path(), None);
6191 let agent = Agent::new(provider, tools, AgentConfig::default());
6192 let session = Arc::new(Mutex::new(Session::in_memory()));
6193 let mut agent_session =
6194 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6195
6196 agent_session
6197 .enable_extensions(&["bash"], temp_dir.path(), None, &[entry_path])
6198 .await
6199 .expect("enable extensions");
6200
6201 let tool_call = ToolCall {
6202 id: "call-1".to_string(),
6203 name: "bash".to_string(),
6204 arguments: json!({ "command": "printf 'hi' > blocked.txt" }),
6205 thought_signature: None,
6206 };
6207
6208 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
6209 let (output, is_error) = agent_session
6210 .agent
6211 .execute_tool(tool_call, on_event, test_turn_latency())
6212 .await;
6213
6214 assert!(is_error);
6215 assert!(output.is_error);
6216 assert_eq!(output.details, None);
6217 assert!(
6218 !temp_dir.path().join("blocked.txt").exists(),
6219 "expected bash command not to run when blocked"
6220 );
6221 assert!(
6222 matches!(output.content.as_slice(), [ContentBlock::Text(_)]),
6223 "Expected text output, got {:?}",
6224 output.content
6225 );
6226 if let [ContentBlock::Text(text)] = output.content.as_slice() {
6227 assert_eq!(text.text, "Tool execution blocked: blocked bash in test");
6228 }
6229 });
6230 }
6231
6232 #[test]
6233 fn tool_result_hook_can_modify_tool_output() {
6234 let runtime = RuntimeBuilder::current_thread()
6235 .build()
6236 .expect("runtime build");
6237
6238 runtime.block_on(async {
6239 let temp_dir = tempfile::tempdir().expect("tempdir");
6240 let entry_path = temp_dir.path().join("ext.mjs");
6241 std::fs::write(
6242 &entry_path,
6243 r#"
6244 export default function init(pi) {
6245 pi.on("tool_result", async (event) => {
6246 if (Object.is(event && event.toolName, "count_tool")) {
6247 return {
6248 content: [{ type: "text", text: "modified" }],
6249 details: { from: "tool_result" }
6250 };
6251 }
6252 return {};
6253 });
6254 }
6255 "#,
6256 )
6257 .expect("write extension entry");
6258
6259 let provider = Arc::new(NoopProvider);
6260 let calls = Arc::new(AtomicUsize::new(0));
6261 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
6262 calls: Arc::clone(&calls),
6263 })]);
6264 let agent = Agent::new(provider, tools, AgentConfig::default());
6265 let session = Arc::new(Mutex::new(Session::in_memory()));
6266 let mut agent_session =
6267 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6268
6269 agent_session
6270 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
6271 .await
6272 .expect("enable extensions");
6273
6274 let tool_call = ToolCall {
6275 id: "call-1".to_string(),
6276 name: "count_tool".to_string(),
6277 arguments: json!({}),
6278 thought_signature: None,
6279 };
6280
6281 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
6282 let (output, is_error) = agent_session
6283 .agent
6284 .execute_tool(tool_call, on_event, test_turn_latency())
6285 .await;
6286
6287 assert!(!is_error);
6288 assert!(!output.is_error);
6289 assert_eq!(calls.load(Ordering::SeqCst), 1);
6290 assert_eq!(output.details, Some(json!({ "from": "tool_result" })));
6291
6292 assert!(
6293 matches!(output.content.as_slice(), [ContentBlock::Text(_)]),
6294 "Expected text output, got {:?}",
6295 output.content
6296 );
6297 if let [ContentBlock::Text(text)] = output.content.as_slice() {
6298 assert_eq!(text.text, "modified");
6299 }
6300 });
6301 }
6302
6303 #[test]
6304 fn tool_result_hook_can_modify_tool_not_found_error() {
6305 let runtime = RuntimeBuilder::current_thread()
6306 .build()
6307 .expect("runtime build");
6308
6309 runtime.block_on(async {
6310 let temp_dir = tempfile::tempdir().expect("tempdir");
6311 let entry_path = temp_dir.path().join("ext.mjs");
6312 std::fs::write(
6313 &entry_path,
6314 r#"
6315 export default function init(pi) {
6316 pi.on("tool_result", async (event) => {
6317 if (Object.is(event && event.toolName, "missing_tool") && event.isError) {
6318 return {
6319 content: [{ type: "text", text: "overridden" }],
6320 details: { handled: true }
6321 };
6322 }
6323 return {};
6324 });
6325 }
6326 "#,
6327 )
6328 .expect("write extension entry");
6329
6330 let provider = Arc::new(NoopProvider);
6331 let tools = ToolRegistry::from_tools(Vec::new());
6332 let agent = Agent::new(provider, tools, AgentConfig::default());
6333 let session = Arc::new(Mutex::new(Session::in_memory()));
6334 let mut agent_session =
6335 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6336
6337 agent_session
6338 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
6339 .await
6340 .expect("enable extensions");
6341
6342 let tool_call = ToolCall {
6343 id: "call-1".to_string(),
6344 name: "missing_tool".to_string(),
6345 arguments: json!({}),
6346 thought_signature: None,
6347 };
6348
6349 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
6350 let (output, is_error) = agent_session
6351 .agent
6352 .execute_tool(tool_call, on_event, test_turn_latency())
6353 .await;
6354
6355 assert!(is_error);
6356 assert!(output.is_error);
6357 assert_eq!(output.details, Some(json!({ "handled": true })));
6358
6359 assert!(
6360 matches!(output.content.as_slice(), [ContentBlock::Text(_)]),
6361 "Expected text output, got {:?}",
6362 output.content
6363 );
6364 if let [ContentBlock::Text(text)] = output.content.as_slice() {
6365 assert_eq!(text.text, "overridden");
6366 }
6367 });
6368 }
6369
6370 #[test]
6371 fn tool_result_hook_errors_fail_open() {
6372 let runtime = RuntimeBuilder::current_thread()
6373 .build()
6374 .expect("runtime build");
6375
6376 runtime.block_on(async {
6377 let temp_dir = tempfile::tempdir().expect("tempdir");
6378 let entry_path = temp_dir.path().join("ext.mjs");
6379 std::fs::write(
6380 &entry_path,
6381 r#"
6382 export default function init(pi) {
6383 pi.on("tool_result", async (_event) => {
6384 throw new Error("boom");
6385 });
6386 }
6387 "#,
6388 )
6389 .expect("write extension entry");
6390
6391 let provider = Arc::new(NoopProvider);
6392 let calls = Arc::new(AtomicUsize::new(0));
6393 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
6394 calls: Arc::clone(&calls),
6395 })]);
6396 let agent = Agent::new(provider, tools, AgentConfig::default());
6397 let session = Arc::new(Mutex::new(Session::in_memory()));
6398 let mut agent_session =
6399 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6400
6401 agent_session
6402 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
6403 .await
6404 .expect("enable extensions");
6405
6406 let tool_call = ToolCall {
6407 id: "call-1".to_string(),
6408 name: "count_tool".to_string(),
6409 arguments: json!({}),
6410 thought_signature: None,
6411 };
6412
6413 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
6414 let (output, is_error) = agent_session
6415 .agent
6416 .execute_tool(tool_call, on_event, test_turn_latency())
6417 .await;
6418
6419 assert!(!is_error);
6420 assert!(!output.is_error);
6421 assert_eq!(calls.load(Ordering::SeqCst), 1);
6422
6423 assert_eq!(output.details, None);
6424 assert!(
6425 matches!(output.content.as_slice(), [ContentBlock::Text(_)]),
6426 "Expected text output, got {:?}",
6427 output.content
6428 );
6429 if let [ContentBlock::Text(text)] = output.content.as_slice() {
6430 assert_eq!(text.text, "ok");
6431 }
6432 });
6433 }
6434
6435 #[test]
6436 fn tool_result_hook_runs_on_blocked_tool_call() {
6437 let runtime = RuntimeBuilder::current_thread()
6438 .build()
6439 .expect("runtime build");
6440
6441 runtime.block_on(async {
6442 let temp_dir = tempfile::tempdir().expect("tempdir");
6443 let entry_path = temp_dir.path().join("ext.mjs");
6444 std::fs::write(
6445 &entry_path,
6446 r#"
6447 export default function init(pi) {
6448 pi.on("tool_call", async (event) => {
6449 if (Object.is(event && event.toolName, "count_tool")) {
6450 return { block: true, reason: "blocked in test" };
6451 }
6452 return {};
6453 });
6454
6455 pi.on("tool_result", async (event) => {
6456 if (Object.is(event && event.toolName, "count_tool") && event.isError) {
6457 return { content: [{ type: "text", text: "override" }] };
6458 }
6459 return {};
6460 });
6461 }
6462 "#,
6463 )
6464 .expect("write extension entry");
6465
6466 let provider = Arc::new(NoopProvider);
6467 let calls = Arc::new(AtomicUsize::new(0));
6468 let tools = ToolRegistry::from_tools(vec![Box::new(CountingTool {
6469 calls: Arc::clone(&calls),
6470 })]);
6471 let agent = Agent::new(provider, tools, AgentConfig::default());
6472 let session = Arc::new(Mutex::new(Session::in_memory()));
6473 let mut agent_session =
6474 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6475
6476 agent_session
6477 .enable_extensions(&[], temp_dir.path(), None, &[entry_path])
6478 .await
6479 .expect("enable extensions");
6480
6481 let tool_call = ToolCall {
6482 id: "call-1".to_string(),
6483 name: "count_tool".to_string(),
6484 arguments: json!({}),
6485 thought_signature: None,
6486 };
6487
6488 let on_event: Arc<dyn Fn(AgentEvent) + Send + Sync> = Arc::new(|_| {});
6489 let (output, is_error) = agent_session
6490 .agent
6491 .execute_tool(tool_call, on_event, test_turn_latency())
6492 .await;
6493
6494 assert!(is_error);
6495 assert!(output.is_error);
6496 assert_eq!(calls.load(Ordering::SeqCst), 0);
6497
6498 assert!(
6499 matches!(output.content.as_slice(), [ContentBlock::Text(_)]),
6500 "Expected text output, got {:?}",
6501 output.content
6502 );
6503 if let [ContentBlock::Text(text)] = output.content.as_slice() {
6504 assert_eq!(text.text, "override");
6505 }
6506 });
6507 }
6508}
6509
6510#[cfg(test)]
6511mod abort_tests {
6512 use super::*;
6513 use crate::session::Session;
6514 use crate::tools::{Tool, ToolOutput, ToolRegistry, ToolUpdate};
6515 use asupersync::runtime::RuntimeBuilder;
6516 use async_trait::async_trait;
6517 use futures::Stream;
6518 use serde_json::json;
6519 use std::path::Path;
6520 use std::pin::Pin;
6521 use std::sync::Mutex as StdMutex;
6522 use std::sync::atomic::AtomicUsize;
6523 use std::task::{Context as TaskContext, Poll};
6524
6525 struct StartThenPending {
6526 start: Option<StreamEvent>,
6527 }
6528
6529 impl Stream for StartThenPending {
6530 type Item = crate::error::Result<StreamEvent>;
6531
6532 fn poll_next(
6533 mut self: Pin<&mut Self>,
6534 _cx: &mut TaskContext<'_>,
6535 ) -> Poll<Option<Self::Item>> {
6536 if let Some(event) = self.start.take() {
6537 return Poll::Ready(Some(Ok(event)));
6538 }
6539 Poll::Pending
6540 }
6541 }
6542
6543 #[derive(Debug)]
6544 struct HangingProvider;
6545
6546 #[async_trait]
6547 #[allow(clippy::unnecessary_literal_bound)]
6548 impl Provider for HangingProvider {
6549 fn name(&self) -> &str {
6550 "test-provider"
6551 }
6552
6553 fn api(&self) -> &str {
6554 "test-api"
6555 }
6556
6557 fn model_id(&self) -> &str {
6558 "test-model"
6559 }
6560
6561 async fn stream(
6562 &self,
6563 _context: &Context<'_>,
6564 _options: &StreamOptions,
6565 ) -> crate::error::Result<
6566 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
6567 > {
6568 let partial = AssistantMessage {
6569 content: Vec::new(),
6570 api: self.api().to_string(),
6571 provider: self.name().to_string(),
6572 model: self.model_id().to_string(),
6573 usage: Usage::default(),
6574 stop_reason: StopReason::Stop,
6575 error_message: None,
6576 timestamp: 0,
6577 };
6578
6579 Ok(Box::pin(StartThenPending {
6580 start: Some(StreamEvent::Start { partial }),
6581 }))
6582 }
6583 }
6584
6585 #[derive(Debug)]
6586 struct CountingProvider {
6587 calls: Arc<std::sync::atomic::AtomicUsize>,
6588 }
6589
6590 #[async_trait]
6591 #[allow(clippy::unnecessary_literal_bound)]
6592 impl Provider for CountingProvider {
6593 fn name(&self) -> &str {
6594 "test-provider"
6595 }
6596
6597 fn api(&self) -> &str {
6598 "test-api"
6599 }
6600
6601 fn model_id(&self) -> &str {
6602 "test-model"
6603 }
6604
6605 async fn stream(
6606 &self,
6607 _context: &Context<'_>,
6608 _options: &StreamOptions,
6609 ) -> crate::error::Result<
6610 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
6611 > {
6612 self.calls.fetch_add(1, Ordering::SeqCst);
6613 Ok(Box::pin(futures::stream::empty()))
6614 }
6615 }
6616
6617 #[derive(Debug)]
6618 struct PhasedProvider {
6619 pending_calls: usize,
6620 calls: AtomicUsize,
6621 }
6622
6623 impl PhasedProvider {
6624 const fn new(pending_calls: usize) -> Self {
6625 Self {
6626 pending_calls,
6627 calls: AtomicUsize::new(0),
6628 }
6629 }
6630
6631 fn base_message() -> AssistantMessage {
6632 AssistantMessage {
6633 content: Vec::new(),
6634 api: "test-api".to_string(),
6635 provider: "test-provider".to_string(),
6636 model: "test-model".to_string(),
6637 usage: Usage::default(),
6638 stop_reason: StopReason::Stop,
6639 error_message: None,
6640 timestamp: 0,
6641 }
6642 }
6643 }
6644
6645 #[async_trait]
6646 #[allow(clippy::unnecessary_literal_bound)]
6647 impl Provider for PhasedProvider {
6648 fn name(&self) -> &str {
6649 "test-provider"
6650 }
6651
6652 fn api(&self) -> &str {
6653 "test-api"
6654 }
6655
6656 fn model_id(&self) -> &str {
6657 "test-model"
6658 }
6659
6660 async fn stream(
6661 &self,
6662 _context: &Context<'_>,
6663 _options: &StreamOptions,
6664 ) -> crate::error::Result<
6665 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
6666 > {
6667 let call = self.calls.fetch_add(1, Ordering::SeqCst);
6668 if call < self.pending_calls {
6669 return Ok(Box::pin(StartThenPending {
6670 start: Some(StreamEvent::Start {
6671 partial: Self::base_message(),
6672 }),
6673 }));
6674 }
6675
6676 let partial = Self::base_message();
6677 let mut done = Self::base_message();
6678 done.content = vec![ContentBlock::Text(TextContent::new(format!(
6679 "resumed-response-{call}"
6680 )))];
6681
6682 Ok(Box::pin(futures::stream::iter(vec![
6683 Ok(StreamEvent::Start { partial }),
6684 Ok(StreamEvent::Done {
6685 reason: StopReason::Stop,
6686 message: done,
6687 }),
6688 ])))
6689 }
6690 }
6691
6692 #[derive(Debug)]
6693 struct ToolCallProvider;
6694
6695 #[async_trait]
6696 #[allow(clippy::unnecessary_literal_bound)]
6697 impl Provider for ToolCallProvider {
6698 fn name(&self) -> &str {
6699 "test-provider"
6700 }
6701
6702 fn api(&self) -> &str {
6703 "test-api"
6704 }
6705
6706 fn model_id(&self) -> &str {
6707 "test-model"
6708 }
6709
6710 async fn stream(
6711 &self,
6712 _context: &Context<'_>,
6713 _options: &StreamOptions,
6714 ) -> crate::error::Result<
6715 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
6716 > {
6717 let message = AssistantMessage {
6718 content: vec![ContentBlock::ToolCall(ToolCall {
6719 id: "call-1".to_string(),
6720 name: "hanging_tool".to_string(),
6721 arguments: json!({}),
6722 thought_signature: None,
6723 })],
6724 api: "test-api".to_string(),
6725 provider: "test-provider".to_string(),
6726 model: "test-model".to_string(),
6727 usage: Usage::default(),
6728 stop_reason: StopReason::ToolUse,
6729 error_message: None,
6730 timestamp: 0,
6731 };
6732
6733 Ok(Box::pin(futures::stream::iter(vec![Ok(
6734 StreamEvent::Done {
6735 reason: StopReason::ToolUse,
6736 message,
6737 },
6738 )])))
6739 }
6740 }
6741
6742 #[derive(Debug)]
6743 struct HangingTool;
6744
6745 #[async_trait]
6746 #[allow(clippy::unnecessary_literal_bound)]
6747 impl Tool for HangingTool {
6748 fn name(&self) -> &str {
6749 "hanging_tool"
6750 }
6751
6752 fn label(&self) -> &str {
6753 "Hanging Tool"
6754 }
6755
6756 fn description(&self) -> &str {
6757 "Never completes unless aborted by the host"
6758 }
6759
6760 fn parameters(&self) -> serde_json::Value {
6761 json!({
6762 "type": "object",
6763 "properties": {},
6764 "additionalProperties": false
6765 })
6766 }
6767
6768 async fn execute(
6769 &self,
6770 _tool_call_id: &str,
6771 _input: serde_json::Value,
6772 _on_update: Option<Box<dyn Fn(ToolUpdate) + Send + Sync>>,
6773 ) -> crate::error::Result<ToolOutput> {
6774 futures::future::pending::<()>().await;
6775 unreachable!("hanging tool should be aborted by the agent")
6776 }
6777 }
6778
6779 fn event_tag(event: &AgentEvent) -> &'static str {
6780 match event {
6781 AgentEvent::AgentStart { .. } => "agent_start",
6782 AgentEvent::AgentEnd { error, .. } => {
6783 if error.as_deref() == Some("Aborted") {
6784 "agent_end_aborted"
6785 } else {
6786 "agent_end"
6787 }
6788 }
6789 AgentEvent::TurnStart { .. } => "turn_start",
6790 AgentEvent::TurnEnd { .. } => "turn_end",
6791 AgentEvent::MessageStart { .. } => "message_start",
6792 AgentEvent::MessageUpdate {
6793 assistant_message_event,
6794 ..
6795 } => match &assistant_message_event {
6796 AssistantMessageEvent::Error {
6797 reason: StopReason::Aborted,
6798 ..
6799 } => "assistant_error_aborted",
6800 AssistantMessageEvent::Done { .. } => "assistant_done",
6801 _ => "assistant_update",
6802 },
6803 AgentEvent::MessageEnd { .. } => "message_end",
6804 AgentEvent::ToolExecutionStart { .. } => "tool_start",
6805 AgentEvent::ToolExecutionUpdate { .. } => "tool_update",
6806 AgentEvent::ToolExecutionEnd { .. } => "tool_end",
6807 AgentEvent::AutoCompactionStart { .. } => "auto_compaction_start",
6808 AgentEvent::AutoCompactionEnd { .. } => "auto_compaction_end",
6809 AgentEvent::AutoRetryStart { .. } => "auto_retry_start",
6810 AgentEvent::AutoRetryEnd { .. } => "auto_retry_end",
6811 AgentEvent::ExtensionError { .. } => "extension_error",
6812 }
6813 }
6814
6815 fn assert_abort_resume_message_sequence(persisted: &[Message]) {
6816 assert_eq!(
6817 persisted.len(),
6818 6,
6819 "expected three user+assistant pairs, got: {persisted:?}"
6820 );
6821
6822 let assistant_states = persisted
6823 .iter()
6824 .filter_map(|message| match message {
6825 Message::Assistant(assistant) => Some(assistant.stop_reason),
6826 _ => None,
6827 })
6828 .collect::<Vec<_>>();
6829 assert_eq!(
6830 assistant_states,
6831 vec![StopReason::Aborted, StopReason::Aborted, StopReason::Stop]
6832 );
6833 }
6834
6835 fn assert_abort_resume_timeline_boundaries(timeline: &[String]) {
6836 assert!(
6837 timeline
6838 .iter()
6839 .any(|event| event.as_str().eq("run0:agent_end_aborted")),
6840 "missing aborted boundary for first run: {timeline:?}"
6841 );
6842 assert!(
6843 timeline
6844 .iter()
6845 .any(|event| event.as_str().eq("run1:agent_end_aborted")),
6846 "missing aborted boundary for second run: {timeline:?}"
6847 );
6848 assert!(
6849 timeline
6850 .iter()
6851 .any(|event| event.as_str().eq("run2:agent_end")),
6852 "missing successful boundary for resumed run: {timeline:?}"
6853 );
6854 }
6855
6856 #[test]
6857 fn abort_interrupts_in_flight_stream() {
6858 let runtime = RuntimeBuilder::current_thread()
6859 .build()
6860 .expect("runtime build");
6861 let handle = runtime.handle();
6862
6863 let started = Arc::new(Notify::new());
6864 let started_wait = started.notified();
6865
6866 let (abort_handle, abort_signal) = AbortHandle::new();
6867
6868 let provider = Arc::new(HangingProvider);
6869 let tools = ToolRegistry::new(&[], Path::new("."), None);
6870 let agent = Agent::new(provider, tools, AgentConfig::default());
6871 let session = Arc::new(Mutex::new(Session::in_memory()));
6872 let mut agent_session =
6873 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6874
6875 let started_tx = Arc::clone(&started);
6876 let join = handle.spawn(async move {
6877 agent_session
6878 .run_text_with_abort("hello".to_string(), Some(abort_signal), move |event| {
6879 if matches!(
6880 event,
6881 AgentEvent::MessageStart {
6882 message: Message::Assistant(_)
6883 }
6884 ) {
6885 started_tx.notify_one();
6886 }
6887 })
6888 .await
6889 });
6890
6891 runtime.block_on(async move {
6892 started_wait.await;
6893 abort_handle.abort();
6894
6895 let message = join.await.expect("run_text_with_abort");
6896 assert_eq!(message.stop_reason, StopReason::Aborted);
6897 assert_eq!(message.error_message.as_deref(), Some("Aborted"));
6898 });
6899 }
6900
6901 #[test]
6902 fn ambient_cancellation_interrupts_in_flight_stream() {
6903 let runtime = RuntimeBuilder::current_thread()
6904 .build()
6905 .expect("runtime build");
6906
6907 runtime.block_on(async move {
6908 let (started_tx, started_rx) = std::sync::mpsc::channel();
6909
6910 let provider = Arc::new(HangingProvider);
6911 let tools = ToolRegistry::new(&[], Path::new("."), None);
6912 let agent = Agent::new(provider, tools, AgentConfig::default());
6913 let session = Arc::new(Mutex::new(Session::in_memory()));
6914 let mut agent_session =
6915 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6916
6917 let ambient_cx = asupersync::Cx::for_testing();
6918 let cancel_cx = ambient_cx.clone();
6919 let _current = asupersync::Cx::set_current(Some(ambient_cx));
6920
6921 let cancel_thread = std::thread::spawn(move || {
6922 started_rx
6923 .recv_timeout(std::time::Duration::from_secs(1))
6924 .expect("stream start");
6925 cancel_cx.set_cancel_requested(true);
6926 });
6927
6928 let run = agent_session.run_text_with_abort("hello".to_string(), None, move |event| {
6929 if matches!(
6930 event,
6931 AgentEvent::MessageStart {
6932 message: Message::Assistant(_)
6933 }
6934 ) {
6935 let _ = started_tx.send(());
6936 }
6937 });
6938 futures::pin_mut!(run);
6939
6940 let message = asupersync::time::timeout(
6941 asupersync::time::wall_now(),
6942 std::time::Duration::from_secs(1),
6943 run,
6944 )
6945 .await
6946 .expect("ambient cancellation should finish before timeout")
6947 .expect("run_text_with_abort");
6948
6949 cancel_thread.join().expect("cancel thread");
6950
6951 assert_eq!(message.stop_reason, StopReason::Aborted);
6952 assert_eq!(message.error_message.as_deref(), Some("Aborted"));
6953 });
6954 }
6955
6956 #[test]
6957 fn abort_before_run_skips_provider_stream_call() {
6958 let runtime = RuntimeBuilder::current_thread()
6959 .build()
6960 .expect("runtime build");
6961
6962 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
6963 let provider = Arc::new(CountingProvider {
6964 calls: Arc::clone(&calls),
6965 });
6966 let tools = ToolRegistry::new(&[], Path::new("."), None);
6967 let agent = Agent::new(provider, tools, AgentConfig::default());
6968 let session = Arc::new(Mutex::new(Session::in_memory()));
6969 let mut agent_session =
6970 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
6971
6972 let (abort_handle, abort_signal) = AbortHandle::new();
6973 abort_handle.abort();
6974
6975 runtime.block_on(async move {
6976 let message = agent_session
6977 .run_text_with_abort("hello".to_string(), Some(abort_signal), |_| {})
6978 .await
6979 .expect("run_text_with_abort");
6980 assert_eq!(message.stop_reason, StopReason::Aborted);
6981 assert_eq!(calls.load(Ordering::SeqCst), 0);
6982 });
6983 }
6984
6985 #[test]
6986 fn abort_then_resume_preserves_session_history() {
6987 let runtime = RuntimeBuilder::current_thread()
6988 .build()
6989 .expect("runtime build");
6990 let handle = runtime.handle();
6991
6992 runtime.block_on(async move {
6993 let provider = Arc::new(PhasedProvider::new(1));
6994 let tools = ToolRegistry::new(&[], Path::new("."), None);
6995 let agent = Agent::new(provider, tools, AgentConfig::default());
6996 let session = Arc::new(Mutex::new(Session::in_memory()));
6997 let mut agent_session = AgentSession::new(
6998 agent,
6999 Arc::clone(&session),
7000 false,
7001 ResolvedCompactionSettings::default(),
7002 );
7003
7004 let started = Arc::new(Notify::new());
7005 let (abort_handle, abort_signal) = AbortHandle::new();
7006 let started_for_abort = Arc::clone(&started);
7007 let abort_join = handle.spawn(async move {
7008 started_for_abort.notified().await;
7009 abort_handle.abort();
7010 });
7011
7012 let aborted = agent_session
7013 .run_text_with_abort("first".to_string(), Some(abort_signal), {
7014 let started = Arc::clone(&started);
7015 move |event| {
7016 if matches!(
7017 event,
7018 AgentEvent::MessageStart {
7019 message: Message::Assistant(_)
7020 }
7021 ) {
7022 started.notify_one();
7023 }
7024 }
7025 })
7026 .await
7027 .expect("first run");
7028 abort_join.await;
7029
7030 assert_eq!(aborted.stop_reason, StopReason::Aborted);
7031 assert_eq!(aborted.error_message.as_deref(), Some("Aborted"));
7032
7033 let resumed = agent_session
7034 .run_text("second".to_string(), |_| {})
7035 .await
7036 .expect("resumed run");
7037 assert_eq!(resumed.stop_reason, StopReason::Stop);
7038 assert!(resumed.error_message.is_none());
7039
7040 let cx = crate::agent_cx::AgentCx::for_request();
7041 let persisted = session
7042 .lock(cx.cx())
7043 .await
7044 .expect("lock session")
7045 .to_messages_for_current_path();
7046
7047 assert_eq!(
7048 persisted.len(),
7049 4,
7050 "unexpected message history after abort+resume: {persisted:?}"
7051 );
7052 assert!(matches!(persisted.first(), Some(Message::User(_))));
7053 assert!(matches!(
7054 persisted.get(1),
7055 Some(Message::Assistant(assistant))
7056 if matches!(assistant.stop_reason, StopReason::Aborted)
7057 ));
7058 assert!(matches!(persisted.get(2), Some(Message::User(_))));
7059 assert!(matches!(
7060 persisted.get(3),
7061 Some(Message::Assistant(assistant))
7062 if matches!(assistant.stop_reason, StopReason::Stop)
7063 && assistant.error_message.is_none()
7064 ));
7065 });
7066 }
7067
7068 #[test]
7069 fn repeated_abort_then_resume_has_consistent_timeline_and_state() {
7070 let runtime = RuntimeBuilder::current_thread()
7071 .build()
7072 .expect("runtime build");
7073 let handle = runtime.handle();
7074
7075 runtime.block_on(async move {
7076 let provider = Arc::new(PhasedProvider::new(2));
7077 let tools = ToolRegistry::new(&[], Path::new("."), None);
7078 let agent = Agent::new(provider, tools, AgentConfig::default());
7079 let session = Arc::new(Mutex::new(Session::in_memory()));
7080 let mut agent_session = AgentSession::new(
7081 agent,
7082 Arc::clone(&session),
7083 false,
7084 ResolvedCompactionSettings::default(),
7085 );
7086
7087 let timeline = Arc::new(StdMutex::new(Vec::<String>::new()));
7088
7089 for run_idx in 0..2 {
7090 let started = Arc::new(Notify::new());
7091 let (abort_handle, abort_signal) = AbortHandle::new();
7092 let started_for_abort = Arc::clone(&started);
7093 let abort_join = handle.spawn(async move {
7094 started_for_abort.notified().await;
7095 abort_handle.abort();
7096 });
7097
7098 let run_timeline = Arc::clone(&timeline);
7099 let aborted = agent_session
7100 .run_text_with_abort(format!("abort-run-{run_idx}"), Some(abort_signal), {
7101 let started = Arc::clone(&started);
7102 move |event| {
7103 if let Ok(mut events) = run_timeline.lock() {
7104 events.push(format!("run{run_idx}:{}", event_tag(&event)));
7105 }
7106 if matches!(
7107 event,
7108 AgentEvent::MessageStart {
7109 message: Message::Assistant(_)
7110 }
7111 ) {
7112 started.notify_one();
7113 }
7114 }
7115 })
7116 .await
7117 .expect("aborted run");
7118 abort_join.await;
7119
7120 assert_eq!(
7121 aborted.stop_reason,
7122 StopReason::Aborted,
7123 "run {run_idx} should abort cleanly"
7124 );
7125 }
7126
7127 let run_timeline = Arc::clone(&timeline);
7128 let resumed = agent_session
7129 .run_text("final-run".to_string(), move |event| {
7130 if let Ok(mut events) = run_timeline.lock() {
7131 events.push(format!("run2:{}", event_tag(&event)));
7132 }
7133 })
7134 .await
7135 .expect("final resumed run");
7136 assert_eq!(resumed.stop_reason, StopReason::Stop);
7137 assert!(resumed.error_message.is_none());
7138
7139 let cx = crate::agent_cx::AgentCx::for_request();
7140 let persisted = session
7141 .lock(cx.cx())
7142 .await
7143 .expect("lock session")
7144 .to_messages_for_current_path();
7145
7146 assert_abort_resume_message_sequence(&persisted);
7147
7148 let timeline = timeline
7149 .lock()
7150 .unwrap_or_else(std::sync::PoisonError::into_inner)
7151 .clone();
7152 assert_abort_resume_timeline_boundaries(&timeline);
7153 });
7154 }
7155
7156 #[test]
7157 fn abort_during_tool_execution_records_aborted_tool_result() {
7158 let runtime = RuntimeBuilder::current_thread()
7159 .build()
7160 .expect("runtime build");
7161 let handle = runtime.handle();
7162
7163 runtime.block_on(async move {
7164 let provider = Arc::new(ToolCallProvider);
7165 let tools = ToolRegistry::from_tools(vec![Box::new(HangingTool)]);
7166 let agent = Agent::new(provider, tools, AgentConfig::default());
7167 let session = Arc::new(Mutex::new(Session::in_memory()));
7168 let mut agent_session = AgentSession::new(
7169 agent,
7170 Arc::clone(&session),
7171 false,
7172 ResolvedCompactionSettings::default(),
7173 );
7174
7175 let tool_started = Arc::new(Notify::new());
7176 let (abort_handle, abort_signal) = AbortHandle::new();
7177 let tool_started_for_abort = Arc::clone(&tool_started);
7178 let abort_join = handle.spawn(async move {
7179 tool_started_for_abort.notified().await;
7180 abort_handle.abort();
7181 });
7182
7183 let result = agent_session
7184 .run_text_with_abort("trigger tool".to_string(), Some(abort_signal), {
7185 let tool_started = Arc::clone(&tool_started);
7186 move |event| {
7187 if matches!(event, AgentEvent::ToolExecutionStart { .. }) {
7188 tool_started.notify_one();
7189 }
7190 }
7191 })
7192 .await
7193 .expect("tool-abort run");
7194 abort_join.await;
7195 assert_eq!(result.stop_reason, StopReason::Aborted);
7196
7197 let cx = crate::agent_cx::AgentCx::for_request();
7198 let persisted = session
7199 .lock(cx.cx())
7200 .await
7201 .expect("lock session")
7202 .to_messages_for_current_path();
7203
7204 let tool_result = persisted
7205 .iter()
7206 .find_map(|message| match message {
7207 Message::ToolResult(result) => Some(result),
7208 _ => None,
7209 })
7210 .expect("expected tool result message");
7211 assert!(tool_result.is_error);
7212 assert!(
7213 tool_result.content.iter().any(|block| {
7214 matches!(
7215 block,
7216 ContentBlock::Text(text) if text.text.contains("Tool execution aborted")
7217 )
7218 }),
7219 "missing aborted tool marker in tool output: {:?}",
7220 tool_result.content
7221 );
7222 let details = tool_result
7223 .details
7224 .as_ref()
7225 .expect("aborted tool result should include structured details");
7226 assert_eq!(details["schema"], TOOL_CANCELLATION_SCHEMA_V1);
7227 assert_eq!(details["status"], "cancelled");
7228 assert_eq!(details["reason"], "abort_signal");
7229 assert_eq!(details["toolName"], "hanging_tool");
7230 assert_eq!(details["cleanup"], "tool_result_recorded_no_success");
7231 });
7232 }
7233}
7234
7235#[cfg(test)]
7236mod turn_event_tests {
7237 use super::*;
7238 use crate::session::Session;
7239 use crate::tools::{Tool, ToolOutput, ToolRegistry, ToolUpdate};
7240 use asupersync::runtime::RuntimeBuilder;
7241 use async_trait::async_trait;
7242 use futures::Stream;
7243 use serde_json::json;
7244 use std::path::Path;
7245 use std::pin::Pin;
7246 use std::sync::atomic::AtomicUsize;
7247 fn assistant_message(text: &str) -> AssistantMessage {
7251 AssistantMessage {
7252 content: vec![ContentBlock::Text(TextContent::new(text))],
7253 api: "test-api".to_string(),
7254 provider: "test-provider".to_string(),
7255 model: "test-model".to_string(),
7256 usage: Usage::default(),
7257 stop_reason: StopReason::Stop,
7258 error_message: None,
7259 timestamp: 0,
7260 }
7261 }
7262
7263 struct SingleShotProvider;
7264
7265 #[async_trait]
7266 #[allow(clippy::unnecessary_literal_bound)]
7267 impl Provider for SingleShotProvider {
7268 fn name(&self) -> &str {
7269 "test-provider"
7270 }
7271
7272 fn api(&self) -> &str {
7273 "test-api"
7274 }
7275
7276 fn model_id(&self) -> &str {
7277 "test-model"
7278 }
7279
7280 async fn stream(
7281 &self,
7282 _context: &Context<'_>,
7283 _options: &StreamOptions,
7284 ) -> crate::error::Result<
7285 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
7286 > {
7287 let partial = assistant_message("");
7288 let final_message = assistant_message("hello");
7289 let events = vec![
7290 Ok(StreamEvent::Start { partial }),
7291 Ok(StreamEvent::Done {
7292 reason: StopReason::Stop,
7293 message: final_message,
7294 }),
7295 ];
7296 Ok(Box::pin(futures::stream::iter(events)))
7297 }
7298 }
7299
7300 struct StreamSetupErrorProvider;
7301
7302 #[async_trait]
7303 #[allow(clippy::unnecessary_literal_bound)]
7304 impl Provider for StreamSetupErrorProvider {
7305 fn name(&self) -> &str {
7306 "test-provider"
7307 }
7308
7309 fn api(&self) -> &str {
7310 "test-api"
7311 }
7312
7313 fn model_id(&self) -> &str {
7314 "test-model"
7315 }
7316
7317 async fn stream(
7318 &self,
7319 _context: &Context<'_>,
7320 _options: &StreamOptions,
7321 ) -> crate::error::Result<
7322 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
7323 > {
7324 Err(Error::api("stream setup failed"))
7325 }
7326 }
7327
7328 #[derive(Debug)]
7329 struct EchoTool;
7330
7331 #[async_trait]
7332 #[allow(clippy::unnecessary_literal_bound)]
7333 impl Tool for EchoTool {
7334 fn name(&self) -> &str {
7335 "echo_tool"
7336 }
7337
7338 fn label(&self) -> &str {
7339 "echo_tool"
7340 }
7341
7342 fn description(&self) -> &str {
7343 "echo test tool"
7344 }
7345
7346 fn parameters(&self) -> serde_json::Value {
7347 json!({ "type": "object" })
7348 }
7349
7350 async fn execute(
7351 &self,
7352 _tool_call_id: &str,
7353 _input: serde_json::Value,
7354 _on_update: Option<Box<dyn Fn(ToolUpdate) + Send + Sync>>,
7355 ) -> Result<ToolOutput> {
7356 Ok(ToolOutput {
7357 content: vec![ContentBlock::Text(TextContent::new("tool-ok"))],
7358 details: None,
7359 is_error: false,
7360 })
7361 }
7362 }
7363
7364 #[derive(Debug)]
7365 struct ToolTurnProvider {
7366 calls: AtomicUsize,
7367 }
7368
7369 impl ToolTurnProvider {
7370 const fn new() -> Self {
7371 Self {
7372 calls: AtomicUsize::new(0),
7373 }
7374 }
7375
7376 fn assistant_message_with(
7377 &self,
7378 stop_reason: StopReason,
7379 content: Vec<ContentBlock>,
7380 ) -> AssistantMessage {
7381 AssistantMessage {
7382 content,
7383 api: self.api().to_string(),
7384 provider: self.name().to_string(),
7385 model: self.model_id().to_string(),
7386 usage: Usage::default(),
7387 stop_reason,
7388 error_message: None,
7389 timestamp: 0,
7390 }
7391 }
7392 }
7393
7394 #[async_trait]
7395 #[allow(clippy::unnecessary_literal_bound)]
7396 impl Provider for ToolTurnProvider {
7397 fn name(&self) -> &str {
7398 "test-provider"
7399 }
7400
7401 fn api(&self) -> &str {
7402 "test-api"
7403 }
7404
7405 fn model_id(&self) -> &str {
7406 "test-model"
7407 }
7408
7409 async fn stream(
7410 &self,
7411 _context: &Context<'_>,
7412 _options: &StreamOptions,
7413 ) -> crate::error::Result<
7414 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
7415 > {
7416 let call_index = self.calls.fetch_add(1, Ordering::SeqCst);
7417 let partial = self.assistant_message_with(StopReason::Stop, Vec::new());
7418 let done = if call_index == 0 {
7419 self.assistant_message_with(
7420 StopReason::ToolUse,
7421 vec![ContentBlock::ToolCall(ToolCall {
7422 id: "tool-1".to_string(),
7423 name: "echo_tool".to_string(),
7424 arguments: json!({}),
7425 thought_signature: None,
7426 })],
7427 )
7428 } else {
7429 self.assistant_message_with(
7430 StopReason::Stop,
7431 vec![ContentBlock::Text(TextContent::new("final"))],
7432 )
7433 };
7434
7435 Ok(Box::pin(futures::stream::iter(vec![
7436 Ok(StreamEvent::Start { partial }),
7437 Ok(StreamEvent::Done {
7438 reason: done.stop_reason,
7439 message: done,
7440 }),
7441 ])))
7442 }
7443 }
7444
7445 #[test]
7446 fn turn_events_wrap_assistant_response() {
7447 let runtime = RuntimeBuilder::current_thread()
7448 .build()
7449 .expect("runtime build");
7450 let handle = runtime.handle();
7451
7452 let provider = Arc::new(SingleShotProvider);
7453 let tools = ToolRegistry::new(&[], Path::new("."), None);
7454 let agent = Agent::new(provider, tools, AgentConfig::default());
7455 let session = Arc::new(Mutex::new(Session::in_memory()));
7456 let mut agent_session =
7457 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
7458
7459 let events: Arc<std::sync::Mutex<Vec<AgentEvent>>> =
7460 Arc::new(std::sync::Mutex::new(Vec::new()));
7461 let events_capture = Arc::clone(&events);
7462
7463 let join = handle.spawn(async move {
7464 agent_session
7465 .run_text("hello".to_string(), move |event| {
7466 events_capture
7467 .lock()
7468 .unwrap_or_else(std::sync::PoisonError::into_inner)
7469 .push(event);
7470 })
7471 .await
7472 .expect("run_text")
7473 });
7474
7475 runtime.block_on(async move {
7476 let message = join.await;
7477 assert_eq!(message.stop_reason, StopReason::Stop);
7478
7479 let events = events
7480 .lock()
7481 .unwrap_or_else(std::sync::PoisonError::into_inner);
7482 let turn_start_indices = events
7483 .iter()
7484 .enumerate()
7485 .filter_map(|(idx, event)| {
7486 matches!(event, AgentEvent::TurnStart { .. }).then_some(idx)
7487 })
7488 .collect::<Vec<_>>();
7489 let turn_end_indices = events
7490 .iter()
7491 .enumerate()
7492 .filter_map(|(idx, event)| {
7493 matches!(event, AgentEvent::TurnEnd { .. }).then_some(idx)
7494 })
7495 .collect::<Vec<_>>();
7496
7497 assert_eq!(turn_start_indices.len(), 1);
7498 assert_eq!(turn_end_indices.len(), 1);
7499 assert!(turn_start_indices[0] < turn_end_indices[0]);
7500
7501 let assistant_message_end = events
7502 .iter()
7503 .enumerate()
7504 .find_map(|(idx, event)| match event {
7505 AgentEvent::MessageEnd {
7506 message: Message::Assistant(_),
7507 } => Some(idx),
7508 _ => None,
7509 })
7510 .expect("assistant message end");
7511
7512 assert!(assistant_message_end < turn_end_indices[0]);
7513
7514 let (message_is_assistant, tool_results_empty) = {
7515 let turn_end_event = &events[turn_end_indices[0]];
7516 assert!(
7517 matches!(turn_end_event, AgentEvent::TurnEnd { .. }),
7518 "Expected TurnEnd event, got {turn_end_event:?}"
7519 );
7520 match turn_end_event {
7521 AgentEvent::TurnEnd {
7522 message,
7523 tool_results,
7524 ..
7525 } => (
7526 matches!(message, Message::Assistant(_)),
7527 tool_results.is_empty(),
7528 ),
7529 _ => (false, false),
7530 }
7531 };
7532 drop(events);
7533 assert!(message_is_assistant);
7534 assert!(tool_results_empty);
7535 });
7536 }
7537
7538 #[test]
7539 fn stream_setup_errors_still_emit_turn_end_before_agent_end() {
7540 let runtime = RuntimeBuilder::current_thread()
7541 .build()
7542 .expect("runtime build");
7543 let handle = runtime.handle();
7544
7545 let provider = Arc::new(StreamSetupErrorProvider);
7546 let tools = ToolRegistry::new(&[], Path::new("."), None);
7547 let agent = Agent::new(provider, tools, AgentConfig::default());
7548 let session = Arc::new(Mutex::new(Session::in_memory()));
7549 let mut agent_session =
7550 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
7551
7552 let events: Arc<std::sync::Mutex<Vec<AgentEvent>>> =
7553 Arc::new(std::sync::Mutex::new(Vec::new()));
7554 let events_capture = Arc::clone(&events);
7555
7556 let join = handle.spawn(async move {
7557 agent_session
7558 .run_text("hello".to_string(), move |event| {
7559 events_capture
7560 .lock()
7561 .unwrap_or_else(std::sync::PoisonError::into_inner)
7562 .push(event);
7563 })
7564 .await
7565 .expect_err("run_text should fail before streaming starts")
7566 });
7567
7568 runtime.block_on(async move {
7569 let err = join.await;
7570 assert!(
7571 err.to_string().contains("stream setup failed"),
7572 "unexpected error: {err}"
7573 );
7574
7575 let events = events
7576 .lock()
7577 .unwrap_or_else(std::sync::PoisonError::into_inner);
7578 let turn_start_idx = events
7579 .iter()
7580 .position(|event| matches!(event, AgentEvent::TurnStart { turn_index: 0, .. }))
7581 .expect("turn start");
7582 let turn_end_idx = events
7583 .iter()
7584 .position(|event| matches!(event, AgentEvent::TurnEnd { turn_index: 0, .. }))
7585 .expect("turn end");
7586 let agent_end_idx = events
7587 .iter()
7588 .position(|event| matches!(event, AgentEvent::AgentEnd { .. }))
7589 .expect("agent end");
7590
7591 assert!(turn_start_idx < turn_end_idx);
7592 assert!(turn_end_idx < agent_end_idx);
7593
7594 let assistant_message_end = events
7595 .iter()
7596 .position(|event| {
7597 matches!(
7598 event,
7599 AgentEvent::MessageEnd {
7600 message: Message::Assistant(_),
7601 }
7602 )
7603 })
7604 .expect("assistant message end");
7605 assert!(assistant_message_end < turn_end_idx);
7606
7607 match &events[turn_end_idx] {
7608 AgentEvent::TurnEnd {
7609 message,
7610 tool_results,
7611 ..
7612 } => {
7613 assert!(tool_results.is_empty());
7614 assert!(
7615 matches!(message, Message::Assistant(_)),
7616 "expected assistant message in TurnEnd, got {message:?}"
7617 );
7618 let Message::Assistant(message) = message else {
7619 return;
7620 };
7621 assert_eq!(message.stop_reason, StopReason::Error);
7622 assert_eq!(
7623 message.error_message.as_deref(),
7624 Some("API error: stream setup failed")
7625 );
7626 assert_eq!(message.api, "test-api");
7627 assert_eq!(message.provider, "test-provider");
7628 assert_eq!(message.model, "test-model");
7629 }
7630 other => {
7631 assert!(matches!(other, AgentEvent::TurnEnd { .. }));
7632 return;
7633 }
7634 }
7635
7636 match &events[agent_end_idx] {
7637 AgentEvent::AgentEnd { error, .. } => {
7638 assert_eq!(error.as_deref(), Some("API error: stream setup failed"));
7639 }
7640 other => {
7641 assert!(matches!(other, AgentEvent::AgentEnd { .. }));
7642 }
7643 }
7644 });
7645 }
7646
7647 #[test]
7648 fn turn_events_include_tool_execution_and_tool_result_messages() {
7649 let runtime = RuntimeBuilder::current_thread()
7650 .build()
7651 .expect("runtime build");
7652 let handle = runtime.handle();
7653
7654 let provider = Arc::new(ToolTurnProvider::new());
7655 let tools = ToolRegistry::from_tools(vec![Box::new(EchoTool)]);
7656 let agent = Agent::new(provider, tools, AgentConfig::default());
7657 let session = Arc::new(Mutex::new(Session::in_memory()));
7658 let mut agent_session =
7659 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
7660
7661 let events: Arc<std::sync::Mutex<Vec<AgentEvent>>> =
7662 Arc::new(std::sync::Mutex::new(Vec::new()));
7663 let events_capture = Arc::clone(&events);
7664
7665 let join = handle.spawn(async move {
7666 agent_session
7667 .run_text("hello".to_string(), move |event| {
7668 events_capture
7669 .lock()
7670 .unwrap_or_else(std::sync::PoisonError::into_inner)
7671 .push(event);
7672 })
7673 .await
7674 .expect("run_text")
7675 });
7676
7677 runtime.block_on(async move {
7678 let message = join.await;
7679 assert_eq!(message.stop_reason, StopReason::Stop);
7680
7681 let events = events
7682 .lock()
7683 .unwrap_or_else(std::sync::PoisonError::into_inner);
7684 let turn_start_count = events
7685 .iter()
7686 .filter(|event| matches!(event, AgentEvent::TurnStart { .. }))
7687 .count();
7688 let turn_end_count = events
7689 .iter()
7690 .filter(|event| matches!(event, AgentEvent::TurnEnd { .. }))
7691 .count();
7692 assert_eq!(
7693 turn_start_count, 2,
7694 "expected one tool turn and one final turn"
7695 );
7696 assert_eq!(
7697 turn_end_count, 2,
7698 "expected one tool turn and one final turn"
7699 );
7700
7701 let tool_start_idx = events
7702 .iter()
7703 .position(|event| matches!(event, AgentEvent::ToolExecutionStart { .. }))
7704 .expect("tool execution start event");
7705 let tool_end_idx = events
7706 .iter()
7707 .position(|event| matches!(event, AgentEvent::ToolExecutionEnd { .. }))
7708 .expect("tool execution end event");
7709 assert!(tool_start_idx < tool_end_idx);
7710
7711 let first_turn_end_idx = events
7712 .iter()
7713 .position(|event| matches!(event, AgentEvent::TurnEnd { turn_index: 0, .. }))
7714 .expect("first turn end");
7715 assert!(
7716 tool_end_idx < first_turn_end_idx,
7717 "tool execution should complete before first turn end"
7718 );
7719
7720 let first_turn_tool_results = events.iter().find_map(|event| match event {
7721 AgentEvent::TurnEnd {
7722 turn_index,
7723 tool_results,
7724 ..
7725 } if turn_index.eq(&0) => Some(tool_results),
7726 _ => None,
7727 });
7728
7729 let first_turn_tool_results =
7730 first_turn_tool_results.expect("expected tool results for first turn");
7731 assert_eq!(first_turn_tool_results.len(), 1);
7732 let first_result = first_turn_tool_results.first().unwrap();
7733 if let Message::ToolResult(tr) = first_result {
7734 assert_eq!(tr.tool_name, "echo_tool");
7735 assert!(!tr.is_error);
7736 } else {
7737 unreachable!("expected Message::ToolResult, got {:?}", first_result);
7738 }
7739 drop(events);
7740 });
7741 }
7742}
7743
7744#[derive(Clone)]
7745struct AgentExtensionSession {
7746 handle: SessionHandle,
7747 is_streaming: Arc<AtomicBool>,
7748 is_compacting: Arc<AtomicBool>,
7749 queue_modes: Arc<StdMutex<ExtensionQueueModeState>>,
7750 auto_compaction_enabled: bool,
7751}
7752
7753impl AgentExtensionSession {
7754 fn current_queue_modes(&self) -> (QueueMode, QueueMode) {
7755 self.queue_modes
7756 .lock()
7757 .map_or((QueueMode::OneAtATime, QueueMode::OneAtATime), |state| {
7758 (state.steering_mode, state.follow_up_mode)
7759 })
7760 }
7761
7762 fn state_fallback(&self) -> Value {
7763 let (steering_mode, follow_up_mode) = self.current_queue_modes();
7764 json!({
7765 "model": null,
7766 "thinkingLevel": "off",
7767 "durabilityMode": "balanced",
7768 "isStreaming": self.is_streaming.load(std::sync::atomic::Ordering::SeqCst),
7769 "isCompacting": self.is_compacting.load(std::sync::atomic::Ordering::SeqCst),
7770 "steeringMode": steering_mode.as_str(),
7771 "followUpMode": follow_up_mode.as_str(),
7772 "sessionFile": null,
7773 "sessionId": "",
7774 "sessionName": null,
7775 "autoCompactionEnabled": self.auto_compaction_enabled,
7776 "messageCount": 0,
7777 "pendingMessageCount": 0,
7778 })
7779 }
7780}
7781
7782#[async_trait]
7783impl crate::extensions::ExtensionSession for AgentExtensionSession {
7784 async fn get_state(&self) -> Value {
7785 let (steering_mode, follow_up_mode) = self.current_queue_modes();
7786 let mut state =
7787 <SessionHandle as crate::extensions::ExtensionSession>::get_state(&self.handle).await;
7788 let Some(object) = state.as_object_mut() else {
7789 return self.state_fallback();
7790 };
7791
7792 object.insert(
7793 "isStreaming".to_string(),
7794 Value::Bool(self.is_streaming.load(std::sync::atomic::Ordering::SeqCst)),
7795 );
7796 object.insert(
7797 "isCompacting".to_string(),
7798 Value::Bool(self.is_compacting.load(std::sync::atomic::Ordering::SeqCst)),
7799 );
7800 object.insert(
7801 "steeringMode".to_string(),
7802 Value::String(steering_mode.as_str().to_string()),
7803 );
7804 object.insert(
7805 "followUpMode".to_string(),
7806 Value::String(follow_up_mode.as_str().to_string()),
7807 );
7808 object.insert(
7809 "autoCompactionEnabled".to_string(),
7810 Value::Bool(self.auto_compaction_enabled),
7811 );
7812
7813 state
7814 }
7815
7816 async fn get_messages(&self) -> Vec<crate::session::SessionMessage> {
7817 <SessionHandle as crate::extensions::ExtensionSession>::get_messages(&self.handle).await
7818 }
7819
7820 async fn get_entries(&self) -> Vec<Value> {
7821 <SessionHandle as crate::extensions::ExtensionSession>::get_entries(&self.handle).await
7822 }
7823
7824 async fn get_branch(&self) -> Vec<Value> {
7825 <SessionHandle as crate::extensions::ExtensionSession>::get_branch(&self.handle).await
7826 }
7827
7828 async fn set_name(&self, name: String) -> crate::error::Result<()> {
7829 <SessionHandle as crate::extensions::ExtensionSession>::set_name(&self.handle, name).await
7830 }
7831
7832 async fn append_message(
7833 &self,
7834 message: crate::session::SessionMessage,
7835 ) -> crate::error::Result<()> {
7836 <SessionHandle as crate::extensions::ExtensionSession>::append_message(
7837 &self.handle,
7838 message,
7839 )
7840 .await
7841 }
7842
7843 async fn append_custom_entry(
7844 &self,
7845 custom_type: String,
7846 data: Option<Value>,
7847 ) -> crate::error::Result<()> {
7848 <SessionHandle as crate::extensions::ExtensionSession>::append_custom_entry(
7849 &self.handle,
7850 custom_type,
7851 data,
7852 )
7853 .await
7854 }
7855
7856 async fn set_model(&self, provider: String, model_id: String) -> crate::error::Result<()> {
7857 <SessionHandle as crate::extensions::ExtensionSession>::set_model(
7858 &self.handle,
7859 provider,
7860 model_id,
7861 )
7862 .await
7863 }
7864
7865 async fn get_model(&self) -> (Option<String>, Option<String>) {
7866 <SessionHandle as crate::extensions::ExtensionSession>::get_model(&self.handle).await
7867 }
7868
7869 async fn set_thinking_level(&self, level: String) -> crate::error::Result<()> {
7870 <SessionHandle as crate::extensions::ExtensionSession>::set_thinking_level(
7871 &self.handle,
7872 level,
7873 )
7874 .await
7875 }
7876
7877 async fn get_thinking_level(&self) -> Option<String> {
7878 <SessionHandle as crate::extensions::ExtensionSession>::get_thinking_level(&self.handle)
7879 .await
7880 }
7881
7882 async fn set_label(
7883 &self,
7884 target_id: String,
7885 label: Option<String>,
7886 ) -> crate::error::Result<()> {
7887 <SessionHandle as crate::extensions::ExtensionSession>::set_label(
7888 &self.handle,
7889 target_id,
7890 label,
7891 )
7892 .await
7893 }
7894}
7895
7896impl AgentSession {
7897 pub const fn runtime_repair_mode_from_policy_mode(mode: RepairPolicyMode) -> RepairMode {
7898 match mode {
7899 RepairPolicyMode::Off => RepairMode::Off,
7900 RepairPolicyMode::Suggest => RepairMode::Suggest,
7901 RepairPolicyMode::AutoSafe => RepairMode::AutoSafe,
7902 RepairPolicyMode::AutoStrict => RepairMode::AutoStrict,
7903 }
7904 }
7905
7906 #[allow(clippy::too_many_arguments)]
7907 async fn start_js_extension_runtime(
7908 stage: &'static str,
7909 cwd: &std::path::Path,
7910 tools: Arc<ToolRegistry>,
7911 manager: ExtensionManager,
7912 policy: ExtensionPolicy,
7913 repair_mode: RepairMode,
7914 memory_limit_bytes: usize,
7915 ) -> Result<ExtensionRuntimeHandle> {
7916 let mut config = PiJsRuntimeConfig {
7917 cwd: cwd.display().to_string(),
7918 repair_mode,
7919 ..PiJsRuntimeConfig::default()
7920 };
7921 config.limits.memory_limit_bytes = Some(memory_limit_bytes).filter(|bytes| *bytes > 0);
7922
7923 let runtime =
7924 JsExtensionRuntimeHandle::start_with_policy(config, tools, manager, policy).await?;
7925 tracing::info!(
7926 event = "pi.extension_runtime.engine_decision",
7927 stage,
7928 requested = "quickjs",
7929 selected = "quickjs",
7930 fallback = false,
7931 "Extension runtime engine selected (legacy JS/TS)"
7932 );
7933 Ok(ExtensionRuntimeHandle::Js(runtime))
7934 }
7935
7936 #[allow(clippy::too_many_arguments)]
7937 async fn start_native_extension_runtime(
7938 stage: &'static str,
7939 _cwd: &std::path::Path,
7940 _tools: Arc<ToolRegistry>,
7941 _manager: ExtensionManager,
7942 _policy: ExtensionPolicy,
7943 _repair_mode: RepairMode,
7944 _memory_limit_bytes: usize,
7945 ) -> Result<ExtensionRuntimeHandle> {
7946 let runtime = NativeRustExtensionRuntimeHandle::start().await?;
7947 tracing::info!(
7948 event = "pi.extension_runtime.engine_decision",
7949 stage,
7950 requested = "native-rust",
7951 selected = "native-rust",
7952 fallback = false,
7953 "Extension runtime engine selected (native-rust)"
7954 );
7955 Ok(ExtensionRuntimeHandle::NativeRust(runtime))
7956 }
7957
7958 pub fn new(
7959 agent: Agent,
7960 session: Arc<Mutex<Session>>,
7961 save_enabled: bool,
7962 compaction_settings: ResolvedCompactionSettings,
7963 ) -> Self {
7964 let extension_ai_completion = Arc::new(StdMutex::new(ExtensionAiCompletionHostState {
7965 provider: agent.provider(),
7966 stream_options: agent.stream_options().clone(),
7967 models: Vec::new(),
7968 }));
7969
7970 Self {
7971 agent,
7972 session,
7973 save_enabled,
7974 input_source: InputSource::Interactive,
7975 extensions: None,
7976 extensions_is_streaming: Arc::new(AtomicBool::new(false)),
7977 extensions_is_compacting: Arc::new(AtomicBool::new(false)),
7978 extensions_turn_active: Arc::new(AtomicBool::new(false)),
7979 extensions_pending_idle_actions: Arc::new(StdMutex::new(VecDeque::new())),
7980 extension_queue_modes: None,
7981 extension_injected_queue: None,
7982 extension_ai_completion,
7983 compaction_settings,
7984 compaction_runtime: None,
7985 runtime_handle: None,
7986 compaction_worker: CompactionWorkerState::new(CompactionQuota::default()),
7987 model_registry: None,
7988 auth_storage: None,
7989 api_key_override: None,
7990 semantic_context_bundle: None,
7991 }
7992 }
7993
7994 pub const fn set_input_source(&mut self, source: InputSource) {
7995 self.input_source = source;
7996 }
7997
7998 #[must_use]
7999 pub fn with_runtime_handle(mut self, runtime_handle: RuntimeHandle) -> Self {
8000 self.compaction_runtime = None;
8001 self.runtime_handle = Some(runtime_handle);
8002 self
8003 }
8004
8005 #[must_use]
8006 pub fn with_model_registry(mut self, registry: ModelRegistry) -> Self {
8007 self.set_model_registry(registry);
8008 self
8009 }
8010
8011 #[must_use]
8012 pub fn with_auth_storage(mut self, auth: AuthStorage) -> Self {
8013 self.auth_storage = Some(auth);
8014 self
8015 }
8016
8017 pub fn set_model_registry(&mut self, registry: ModelRegistry) {
8018 self.set_extension_ai_models(pi_ai_model_registry_values(®istry));
8019 self.model_registry = Some(registry);
8020 }
8021
8022 pub fn set_auth_storage(&mut self, auth: AuthStorage) {
8023 self.auth_storage = Some(auth);
8024 }
8025
8026 #[must_use]
8027 pub fn with_api_key_override(mut self, api_key: Option<String>) -> Self {
8028 self.set_api_key_override(api_key);
8029 self
8030 }
8031
8032 pub fn set_api_key_override(&mut self, api_key: Option<String>) {
8033 self.api_key_override = normalize_api_key_opt(api_key);
8034 }
8035
8036 pub fn refresh_extension_completion_host_state(&self) {
8037 let Ok(mut state) = self.extension_ai_completion.lock() else {
8038 tracing::error!("extension completion host state mutex poisoned; keeping stale state");
8039 return;
8040 };
8041 state.provider = self.agent.provider();
8042 state.stream_options = self.agent.stream_options().clone();
8043 }
8044
8045 fn set_extension_ai_models(&self, models: Vec<Value>) {
8046 let Ok(mut state) = self.extension_ai_completion.lock() else {
8047 tracing::error!(
8048 "extension completion host state mutex poisoned; keeping stale model catalog"
8049 );
8050 return;
8051 };
8052 state.models = models;
8053 }
8054
8055 pub fn set_semantic_context_bundle(
8056 &mut self,
8057 injection: Option<SemanticContextBundleInjection>,
8058 ) {
8059 self.semantic_context_bundle = injection;
8060 }
8061
8062 pub const fn semantic_context_bundle(&self) -> Option<&SemanticContextBundleInjection> {
8063 self.semantic_context_bundle.as_ref()
8064 }
8065
8066 pub fn set_queue_modes(&mut self, steering_mode: QueueMode, follow_up_mode: QueueMode) {
8067 self.agent.set_queue_modes(steering_mode, follow_up_mode);
8068
8069 if let Some(queue_modes) = &self.extension_queue_modes
8070 && let Ok(mut state) = queue_modes.lock()
8071 {
8072 state.set_modes(steering_mode, follow_up_mode);
8073 }
8074
8075 if let Some(injected_queue) = &self.extension_injected_queue
8076 && let Ok(mut queue) = injected_queue.lock()
8077 {
8078 queue.set_modes(steering_mode, follow_up_mode);
8079 }
8080 }
8081
8082 pub const fn set_compaction_context_window(&mut self, context_window_tokens: u32) {
8083 self.compaction_settings.context_window_tokens = context_window_tokens;
8084 }
8085
8086 pub async fn set_provider_model(&mut self, provider_id: &str, model_id: &str) -> Result<()> {
8087 let already_active = {
8088 let provider = self.agent.provider();
8089 provider.name().eq(provider_id) && provider.model_id().eq(model_id)
8090 };
8091 let current_thinking = self
8092 .agent
8093 .stream_options()
8094 .thinking_level
8095 .unwrap_or_default();
8096
8097 let target_entry = self
8098 .model_registry
8099 .as_ref()
8100 .and_then(|registry| registry.find(provider_id, model_id));
8101 let next_thinking = if let Some(target_entry) = target_entry {
8102 let resolved_key = self.resolve_stream_api_key_for_model(&target_entry);
8103 if !already_active
8104 && model_requires_configured_credential(&target_entry)
8105 && resolved_key.is_none()
8106 {
8107 return Err(Error::auth(format!(
8108 "Missing credentials for {provider_id}/{model_id}"
8109 )));
8110 }
8111 self.clamp_thinking_level_for_model(provider_id, model_id, current_thinking)
8112 } else if already_active {
8113 current_thinking
8114 } else {
8115 return Err(Error::validation(format!(
8116 "Unable to switch provider/model to {provider_id}/{model_id}"
8117 )));
8118 };
8119
8120 if !already_active {
8121 self.apply_session_model_selection(provider_id, model_id)?;
8122 }
8123 self.agent.stream_options_mut().thinking_level = Some(next_thinking);
8124 self.refresh_extension_completion_host_state();
8125
8126 {
8127 let cx = crate::agent_cx::AgentCx::for_request();
8128 let mut session = self
8129 .session
8130 .lock(cx.cx())
8131 .await
8132 .map_err(|e| Error::session(e.to_string()))?;
8133 let previous_model = session.effective_model_for_current_path();
8134 let previous_thinking = session
8135 .effective_thinking_level_for_current_path()
8136 .as_deref()
8137 .and_then(|value| value.parse::<crate::model::ThinkingLevel>().ok());
8138 if previous_model
8139 .as_ref()
8140 .map(|(provider, model_id)| (provider.as_str(), model_id.as_str()))
8141 != Some((provider_id, model_id))
8142 {
8143 session.append_model_change(provider_id.to_string(), model_id.to_string());
8144 }
8145 session.set_model_header(
8146 Some(provider_id.to_string()),
8147 Some(model_id.to_string()),
8148 Some(next_thinking.to_string()),
8149 );
8150 if !previous_thinking.is_some_and(|previous| previous.eq(&next_thinking)) {
8151 session.append_thinking_level_change(next_thinking.to_string());
8152 }
8153 }
8154
8155 self.persist_session().await
8156 }
8157
8158 pub async fn set_thinking_level(&mut self, level: crate::model::ThinkingLevel) -> Result<()> {
8167 let cx = crate::agent_cx::AgentCx::for_request();
8168 let (effective_level, changed) = {
8169 let mut guard = self
8170 .session
8171 .lock(cx.cx())
8172 .await
8173 .map_err(|e| Error::session(e.to_string()))?;
8174 let (provider_id, model_id) =
8175 guard.effective_model_for_current_path().unwrap_or_else(|| {
8176 let provider = self.agent.provider();
8177 (provider.name().to_string(), provider.model_id().to_string())
8178 });
8179 let effective_level =
8180 self.clamp_thinking_level_for_model(&provider_id, &model_id, level);
8181 let level_string = effective_level.to_string();
8182 let changed = guard.effective_thinking_level_for_current_path().as_deref()
8183 != Some(level_string.as_str());
8184 guard.set_model_header(None, None, Some(level_string.clone()));
8185 if changed {
8186 guard.append_thinking_level_change(level_string);
8187 }
8188 (effective_level, changed)
8189 };
8190 self.agent.stream_options_mut().thinking_level = Some(effective_level);
8191 self.refresh_extension_completion_host_state();
8192 if changed {
8193 self.persist_session().await
8194 } else {
8195 Ok(())
8196 }
8197 }
8198
8199 pub(crate) fn clamp_thinking_level_for_model(
8200 &self,
8201 provider_id: &str,
8202 model_id: &str,
8203 level: crate::model::ThinkingLevel,
8204 ) -> crate::model::ThinkingLevel {
8205 self.model_registry
8206 .as_ref()
8207 .and_then(|registry| registry.find(provider_id, model_id))
8208 .map_or(level, |entry| entry.clamp_thinking_level(level))
8209 }
8210
8211 fn resolve_stream_api_key_for_model(&self, entry: &ModelEntry) -> Option<String> {
8212 let normalize = |key_opt: Option<String>| {
8213 key_opt.and_then(|key| {
8214 let trimmed = key.trim();
8215 (!trimmed.is_empty()).then(|| trimmed.to_string())
8216 })
8217 };
8218
8219 normalize(self.api_key_override.clone())
8220 .or_else(|| {
8221 self.auth_storage
8222 .as_ref()
8223 .and_then(|auth| normalize(auth.resolve_api_key(&entry.model.provider, None)))
8224 })
8225 .or_else(|| normalize(entry.api_key.clone()))
8226 }
8227
8228 pub(crate) async fn sync_runtime_selection_from_session_header(&mut self) -> Result<()> {
8229 let session_state = {
8230 let cx = crate::agent_cx::AgentCx::for_request();
8231 let session = self
8232 .session
8233 .lock(cx.cx())
8234 .await
8235 .map_err(|e| Error::session(e.to_string()))?;
8236 (
8237 session.effective_model_for_current_path(),
8238 session.effective_thinking_level_for_current_path(),
8239 )
8240 };
8241
8242 let (session_model, session_thinking) = session_state;
8243 let current_thinking = self
8244 .agent
8245 .stream_options()
8246 .thinking_level
8247 .unwrap_or_default();
8248
8249 if let Some((provider_id, model_id)) = session_model.as_ref() {
8250 self.apply_session_model_selection(provider_id, model_id)?;
8251 }
8252
8253 let parsed_session_thinking = session_thinking.as_deref().and_then(|raw| {
8254 raw.parse::<crate::model::ThinkingLevel>().map_or_else(
8255 |_| {
8256 tracing::warn!("Ignoring invalid session thinking level: {raw}");
8257 None
8258 },
8259 Some,
8260 )
8261 });
8262 let requested = parsed_session_thinking.unwrap_or(current_thinking);
8263
8264 let effective = if let Some((provider_id, model_id)) = session_model.as_ref() {
8265 self.clamp_thinking_level_for_model(provider_id, model_id, requested)
8266 } else {
8267 requested
8268 };
8269
8270 self.agent.stream_options_mut().thinking_level = Some(effective);
8271 self.refresh_extension_completion_host_state();
8272
8273 let thinking_changed = !effective.eq(¤t_thinking);
8274 let persist_needed = if session_thinking.is_some() {
8275 !parsed_session_thinking.is_some_and(|parsed| parsed.eq(&effective))
8276 } else {
8277 thinking_changed
8278 };
8279 if !persist_needed {
8280 return Ok(());
8281 }
8282
8283 {
8284 let cx = crate::agent_cx::AgentCx::for_request();
8285 let mut session = self
8286 .session
8287 .lock(cx.cx())
8288 .await
8289 .map_err(|e| Error::session(e.to_string()))?;
8290 let previous_thinking = session
8291 .header
8292 .thinking_level
8293 .as_deref()
8294 .and_then(|value| value.parse::<crate::model::ThinkingLevel>().ok());
8295 session.set_model_header(None, None, Some(effective.to_string()));
8296 if thinking_changed
8297 && !previous_thinking.is_some_and(|previous| previous.eq(&effective))
8298 {
8299 session.append_thinking_level_change(effective.to_string());
8300 }
8301 }
8302
8303 self.persist_session().await
8304 }
8305
8306 fn apply_session_model_selection(&mut self, provider_id: &str, model_id: &str) -> Result<()> {
8307 if self.agent.provider().name().eq(provider_id)
8308 && self.agent.provider().model_id().eq(model_id)
8309 {
8310 return Ok(());
8311 }
8312
8313 let Some(registry) = &self.model_registry else {
8314 return Err(Error::validation(format!(
8315 "Unable to switch provider/model to {provider_id}/{model_id}"
8316 )));
8317 };
8318
8319 let Some(entry) = registry.find(provider_id, model_id) else {
8320 return Err(Error::validation(format!(
8321 "Unable to switch provider/model to {provider_id}/{model_id}"
8322 )));
8323 };
8324
8325 let resolved_key = self.resolve_stream_api_key_for_model(&entry);
8326 if model_requires_configured_credential(&entry) && resolved_key.is_none() {
8327 return Err(Error::auth(format!(
8328 "Missing credentials for {provider_id}/{model_id}"
8329 )));
8330 }
8331
8332 match crate::providers::create_provider(
8333 &entry,
8334 self.extensions.as_ref().map(ExtensionRegion::manager),
8335 ) {
8336 Ok(provider) => {
8337 tracing::info!("Updating agent provider to {provider_id}/{model_id}");
8338 self.agent.set_provider(provider);
8339
8340 let stream_options = self.agent.stream_options_mut();
8341 stream_options.api_key.clone_from(&resolved_key);
8342 stream_options.headers.clone_from(&entry.headers);
8343 stream_options.max_tokens = Some(entry.model.max_tokens);
8348 self.refresh_extension_completion_host_state();
8349 Ok(())
8350 }
8351 Err(e) => Err(Error::validation(format!(
8352 "Unable to switch provider/model to {provider_id}/{model_id}: {e}"
8353 ))),
8354 }
8355 }
8356
8357 pub const fn save_enabled(&self) -> bool {
8358 self.save_enabled
8359 }
8360
8361 pub async fn compact_now(
8363 &mut self,
8364 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
8365 ) -> Result<()> {
8366 self.compact_synchronous(Arc::new(on_event)).await
8367 }
8368
8369 pub async fn execute_extension_command(
8370 &mut self,
8371 command_name: &str,
8372 args: &str,
8373 timeout_ms: u64,
8374 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
8375 ) -> Result<Value> {
8376 self.execute_extension_command_with_abort(command_name, args, timeout_ms, None, on_event)
8377 .await
8378 }
8379
8380 pub async fn execute_extension_command_with_abort(
8381 &mut self,
8382 command_name: &str,
8383 args: &str,
8384 timeout_ms: u64,
8385 abort: Option<AbortSignal>,
8386 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
8387 ) -> Result<Value> {
8388 let manager = self
8389 .extensions
8390 .as_ref()
8391 .map(ExtensionRegion::manager)
8392 .ok_or_else(|| Error::extension("Extensions are disabled"))?
8393 .clone();
8394 let on_event: AgentEventHandler = Arc::new(on_event);
8395
8396 self.run_pending_idle_actions_with_abort(abort.clone(), Arc::clone(&on_event))
8397 .await?;
8398
8399 let command_result = manager
8400 .execute_command(command_name, args, timeout_ms)
8401 .await;
8402 let replay_result = self
8403 .run_pending_idle_actions_with_abort(abort, Arc::clone(&on_event))
8404 .await;
8405
8406 match command_result {
8407 Ok(value) => {
8408 replay_result?;
8409 Ok(value)
8410 }
8411 Err(err) => {
8412 if let Err(replay_err) = replay_result {
8413 tracing::warn!(
8414 "extension command follow-up replay failed after command error: {replay_err}"
8415 );
8416 }
8417 Err(err)
8418 }
8419 }
8420 }
8421
8422 #[allow(clippy::too_many_lines)]
8428 async fn maybe_compact(&mut self, on_event: AgentEventHandler) -> Result<()> {
8429 if !self.compaction_settings.enabled {
8430 return Ok(());
8431 }
8432
8433 if let Some(outcome) = self.compaction_worker.try_recv().await {
8435 self.extensions_is_compacting
8436 .store(false, std::sync::atomic::Ordering::SeqCst);
8437 match outcome {
8438 Ok(result) => {
8439 self.apply_compaction_result(result, Arc::clone(&on_event))
8440 .await?;
8441 }
8442 Err(e) => {
8443 on_event(AgentEvent::AutoCompactionEnd {
8444 result: None,
8445 aborted: false,
8446 will_retry: false,
8447 error_message: Some(e.to_string()),
8448 });
8449 }
8450 }
8451 }
8452
8453 if !self.compaction_worker.can_start() {
8455 return Ok(());
8456 }
8457
8458 let (entries, preparation) = {
8459 let cx = crate::agent_cx::AgentCx::for_request();
8460 let mut session = self
8461 .session
8462 .lock(cx.cx())
8463 .await
8464 .map_err(|e| Error::session(e.to_string()))?;
8465 session.ensure_entry_ids();
8466 let entries = session
8467 .entries_for_current_path()
8468 .into_iter()
8469 .cloned()
8470 .collect::<Vec<_>>();
8471 let prep = compaction::prepare_compaction(&entries, self.compaction_settings.clone());
8472 (entries, prep)
8473 };
8474
8475 if let Some(prep) = preparation {
8476 let admission = self
8477 .compaction_worker
8478 .admission_decision(Some(&prep), &CompactionAdmissionSignals::default());
8479 if !admission.allowed {
8480 tracing::info!(
8481 reason = admission.reason.as_str(),
8482 tokens_before = admission.tokens_before,
8483 "Background compaction admission denied"
8484 );
8485 return Ok(());
8486 }
8487
8488 on_event(AgentEvent::AutoCompactionStart {
8489 reason: format!("threshold;admission={}", admission.reason.as_str()),
8490 });
8491
8492 let before_outcome = self.dispatch_before_compact(&prep, &entries, None).await;
8493 if before_outcome.cancel {
8494 on_event(AgentEvent::AutoCompactionEnd {
8495 result: None,
8496 aborted: true,
8497 will_retry: false,
8498 error_message: None,
8499 });
8500 return Ok(());
8501 }
8502
8503 if let Some(compaction) = before_outcome.compaction {
8504 let result_value = Some(Self::auto_compaction_result_payload(
8505 compaction.summary.clone(),
8506 compaction.first_kept_entry_id.clone(),
8507 compaction.tokens_before,
8508 compaction.details.clone(),
8509 ));
8510 self.extensions_is_compacting
8511 .store(true, std::sync::atomic::Ordering::SeqCst);
8512 let apply_result = self
8513 .apply_compaction_entry(
8514 compaction.summary,
8515 compaction.first_kept_entry_id,
8516 compaction.tokens_before,
8517 compaction.details,
8518 true,
8519 )
8520 .await;
8521 self.extensions_is_compacting
8522 .store(false, std::sync::atomic::Ordering::SeqCst);
8523 apply_result?;
8524 on_event(AgentEvent::AutoCompactionEnd {
8525 result: result_value,
8526 aborted: false,
8527 will_retry: false,
8528 error_message: None,
8529 });
8530 return Ok(());
8531 }
8532
8533 let provider = self.agent.provider();
8534 let credential = self
8535 .agent
8536 .stream_options()
8537 .api_key
8538 .clone()
8539 .unwrap_or_default();
8540
8541 let runtime_handle = match self.compaction_runtime_handle() {
8542 Ok(runtime_handle) => runtime_handle,
8543 Err(e) => {
8544 on_event(AgentEvent::AutoCompactionEnd {
8545 result: None,
8546 aborted: false,
8547 will_retry: false,
8548 error_message: Some(e.to_string()),
8549 });
8550 return Ok(());
8551 }
8552 };
8553
8554 self.compaction_worker
8555 .start(&runtime_handle, prep, provider, credential, None);
8556 self.extensions_is_compacting
8557 .store(true, std::sync::atomic::Ordering::SeqCst);
8558 }
8559
8560 Ok(())
8561 }
8562
8563 fn compaction_runtime_handle(&mut self) -> Result<RuntimeHandle> {
8564 if let Some(runtime_handle) = self.runtime_handle.clone() {
8565 return Ok(runtime_handle);
8566 }
8567
8568 let runtime = RuntimeBuilder::new().build().map_err(|e| {
8569 Error::session(format!("Background compaction runtime init failed: {e}"))
8570 })?;
8571 let runtime_handle = runtime.handle();
8572 self.compaction_runtime = Some(runtime);
8573 self.runtime_handle = Some(runtime_handle.clone());
8574 Ok(runtime_handle)
8575 }
8576
8577 fn auto_compaction_result_payload(
8578 summary: String,
8579 first_kept_entry_id: String,
8580 tokens_before: u64,
8581 details: Option<Value>,
8582 ) -> Value {
8583 let mut payload = serde_json::Map::new();
8584 payload.insert("summary".to_string(), Value::String(summary));
8585 payload.insert(
8586 "firstKeptEntryId".to_string(),
8587 Value::String(first_kept_entry_id),
8588 );
8589 payload.insert("tokensBefore".to_string(), Value::from(tokens_before));
8590 if let Some(details) = details {
8591 payload.insert("details".to_string(), details);
8592 }
8593 Value::Object(payload)
8594 }
8595
8596 async fn apply_compaction_entry(
8597 &self,
8598 summary: String,
8599 first_kept_entry_id: String,
8600 tokens_before: u64,
8601 details: Option<Value>,
8602 from_extension: bool,
8603 ) -> Result<()> {
8604 let cx = crate::agent_cx::AgentCx::for_request();
8605 let mut session = OwnedMutexGuard::lock(Arc::clone(&self.session), cx.cx())
8606 .await
8607 .map_err(|e| Error::session(e.to_string()))?;
8608
8609 let from_hook = if from_extension { Some(true) } else { None };
8610 let entry_id = session.append_compaction(
8611 summary,
8612 first_kept_entry_id,
8613 tokens_before,
8614 details,
8615 from_hook,
8616 );
8617
8618 if self.save_enabled {
8619 session
8620 .flush_autosave(AutosaveFlushTrigger::Periodic)
8621 .await?;
8622 }
8623
8624 let compaction_entry = session.get_entry(&entry_id).and_then(|entry| {
8625 if let crate::session::SessionEntry::Compaction(compaction) = entry {
8626 Some(compaction.clone())
8627 } else {
8628 None
8629 }
8630 });
8631 drop(session);
8632
8633 if let (Some(region), Some(compaction_entry)) = (&self.extensions, compaction_entry) {
8634 let payload = json!({
8635 "compactionEntry": compaction_entry,
8636 "fromExtension": from_extension,
8637 });
8638 if let Err(err) = region
8639 .manager()
8640 .dispatch_event(ExtensionEventName::SessionCompact, Some(payload))
8641 .await
8642 {
8643 tracing::warn!("session_compact extension hook failed (fail-open): {err}");
8644 }
8645 }
8646
8647 Ok(())
8648 }
8649
8650 async fn apply_compaction_result(
8652 &self,
8653 result: compaction::CompactionResult,
8654 on_event: AgentEventHandler,
8655 ) -> Result<()> {
8656 let details = Some(compaction::compaction_details_to_value(&result.details)?);
8657 let result_value = Some(Self::auto_compaction_result_payload(
8658 result.summary.clone(),
8659 result.first_kept_entry_id.clone(),
8660 result.tokens_before,
8661 details.clone(),
8662 ));
8663
8664 self.apply_compaction_entry(
8665 result.summary,
8666 result.first_kept_entry_id,
8667 result.tokens_before,
8668 details,
8669 false,
8670 )
8671 .await?;
8672
8673 on_event(AgentEvent::AutoCompactionEnd {
8674 result: result_value,
8675 aborted: false,
8676 will_retry: false,
8677 error_message: None,
8678 });
8679
8680 Ok(())
8681 }
8682
8683 async fn compact_synchronous(&self, on_event: AgentEventHandler) -> Result<()> {
8685 if !self.compaction_settings.enabled {
8686 return Ok(());
8687 }
8688
8689 let (entries, preparation) = {
8690 let cx = crate::agent_cx::AgentCx::for_request();
8691 let mut session = self
8692 .session
8693 .lock(cx.cx())
8694 .await
8695 .map_err(|e| Error::session(e.to_string()))?;
8696 session.ensure_entry_ids();
8697 let entries = session
8698 .entries_for_current_path()
8699 .into_iter()
8700 .cloned()
8701 .collect::<Vec<_>>();
8702 let prep = compaction::prepare_compaction(&entries, self.compaction_settings.clone());
8703 (entries, prep)
8704 };
8705
8706 if let Some(prep) = preparation {
8707 on_event(AgentEvent::AutoCompactionStart {
8708 reason: "threshold".to_string(),
8709 });
8710
8711 let before_outcome = self.dispatch_before_compact(&prep, &entries, None).await;
8712 if before_outcome.cancel {
8713 on_event(AgentEvent::AutoCompactionEnd {
8714 result: None,
8715 aborted: true,
8716 will_retry: false,
8717 error_message: None,
8718 });
8719 return Err(Error::extension("Compaction cancelled".to_string()));
8720 }
8721
8722 if let Some(compaction) = before_outcome.compaction {
8723 let result_value = Some(Self::auto_compaction_result_payload(
8724 compaction.summary.clone(),
8725 compaction.first_kept_entry_id.clone(),
8726 compaction.tokens_before,
8727 compaction.details.clone(),
8728 ));
8729 self.extensions_is_compacting
8730 .store(true, std::sync::atomic::Ordering::SeqCst);
8731 let apply_result = self
8732 .apply_compaction_entry(
8733 compaction.summary,
8734 compaction.first_kept_entry_id,
8735 compaction.tokens_before,
8736 compaction.details,
8737 true,
8738 )
8739 .await;
8740 self.extensions_is_compacting
8741 .store(false, std::sync::atomic::Ordering::SeqCst);
8742 apply_result?;
8743 on_event(AgentEvent::AutoCompactionEnd {
8744 result: result_value,
8745 aborted: false,
8746 will_retry: false,
8747 error_message: None,
8748 });
8749 return Ok(());
8750 }
8751 self.extensions_is_compacting
8752 .store(true, std::sync::atomic::Ordering::SeqCst);
8753
8754 let provider = self.agent.provider();
8755 let credential = self
8756 .agent
8757 .stream_options()
8758 .api_key
8759 .clone()
8760 .unwrap_or_default();
8761
8762 let compaction_result = compaction::compact(prep, provider, &credential, None).await;
8763 self.extensions_is_compacting
8764 .store(false, std::sync::atomic::Ordering::SeqCst);
8765
8766 match compaction_result {
8767 Ok(result) => {
8768 self.apply_compaction_result(result, Arc::clone(&on_event))
8769 .await?;
8770 }
8771 Err(e) => {
8772 on_event(AgentEvent::AutoCompactionEnd {
8773 result: None,
8774 aborted: false,
8775 will_retry: false,
8776 error_message: Some(e.to_string()),
8777 });
8778 return Err(e);
8779 }
8780 }
8781 }
8782 Ok(())
8783 }
8784
8785 fn resolve_extension_policy_for_enable(
8786 config: Option<&crate::config::Config>,
8787 policy: Option<ExtensionPolicy>,
8788 ) -> ExtensionPolicy {
8789 policy.unwrap_or_else(|| {
8790 config.map_or_else(
8791 || crate::config::Config::default().resolve_extension_policy(None),
8792 |cfg| cfg.resolve_extension_policy(None),
8793 )
8794 })
8795 }
8796
8797 pub async fn enable_extensions(
8798 &mut self,
8799 enabled_tools: &[&str],
8800 cwd: &std::path::Path,
8801 config: Option<&crate::config::Config>,
8802 extension_entries: &[std::path::PathBuf],
8803 ) -> Result<()> {
8804 self.enable_extensions_with_policy(
8805 enabled_tools,
8806 cwd,
8807 config,
8808 extension_entries,
8809 None,
8810 None,
8811 None,
8812 )
8813 .await
8814 }
8815
8816 #[allow(clippy::too_many_lines, clippy::too_many_arguments)]
8817 pub async fn enable_extensions_with_policy(
8818 &mut self,
8819 enabled_tools: &[&str],
8820 cwd: &std::path::Path,
8821 config: Option<&crate::config::Config>,
8822 extension_entries: &[std::path::PathBuf],
8823 policy: Option<ExtensionPolicy>,
8824 repair_policy: Option<RepairPolicyMode>,
8825 pre_warmed: Option<PreWarmedExtensionRuntime>,
8826 ) -> Result<()> {
8827 let mut js_specs: Vec<JsExtensionLoadSpec> = Vec::new();
8828 let mut native_specs: Vec<NativeRustExtensionLoadSpec> = Vec::new();
8829 #[cfg(feature = "wasm-host")]
8830 let mut wasm_specs: Vec<WasmExtensionLoadSpec> = Vec::new();
8831
8832 for entry in extension_entries {
8833 match resolve_extension_load_spec(entry)? {
8834 ExtensionLoadSpec::Js(spec) => js_specs.push(spec),
8835 ExtensionLoadSpec::NativeRust(spec) => native_specs.push(spec),
8836 #[cfg(feature = "wasm-host")]
8837 ExtensionLoadSpec::Wasm(spec) => wasm_specs.push(spec),
8838 }
8839 }
8840
8841 if !js_specs.is_empty() && !native_specs.is_empty() {
8842 return Err(Error::validation(
8843 "Mixed extension runtimes are not supported in one session yet. Use either JS/TS extensions (QuickJS) or native-rust descriptors (*.native.json), but not both at once."
8844 .to_string(),
8845 ));
8846 }
8847
8848 #[cfg(feature = "wasm-host")]
8849 if js_specs.is_empty() && native_specs.is_empty() && wasm_specs.is_empty() {
8850 self.extensions = None;
8851 self.agent.extensions = None;
8852 self.extension_queue_modes = None;
8853 self.extension_injected_queue = None;
8854 return Ok(());
8855 }
8856
8857 #[cfg(not(feature = "wasm-host"))]
8858 if js_specs.is_empty() && native_specs.is_empty() {
8859 self.extensions = None;
8860 self.agent.extensions = None;
8861 self.extension_queue_modes = None;
8862 self.extension_injected_queue = None;
8863 return Ok(());
8864 }
8865
8866 let resolved_policy = Self::resolve_extension_policy_for_enable(config, policy);
8867 let resolved_repair_policy = repair_policy
8868 .or_else(|| config.map(|cfg| cfg.resolve_repair_policy(None)))
8869 .unwrap_or(RepairPolicyMode::AutoSafe);
8870 let runtime_repair_mode =
8871 Self::runtime_repair_mode_from_policy_mode(resolved_repair_policy);
8872 let memory_limit_bytes =
8873 (resolved_policy.max_memory_mb as usize).saturating_mul(1024 * 1024);
8874 let wants_js_runtime = !js_specs.is_empty();
8875
8876 #[allow(unused_variables)]
8879 let (manager, tools) = if let Some(pre) = pre_warmed {
8880 let manager = pre.manager;
8881 let tools = pre.tools;
8882 let runtime = match pre.runtime {
8883 ExtensionRuntimeHandle::NativeRust(runtime) => {
8884 if wants_js_runtime {
8885 tracing::warn!(
8886 event = "pi.extension_runtime.prewarm.mismatch",
8887 expected = "quickjs",
8888 got = "native-rust",
8889 "Pre-warmed runtime mismatched requested JS mode; creating quickjs runtime"
8890 );
8891 Self::start_js_extension_runtime(
8892 "agent_enable_extensions_prewarm_mismatch",
8893 cwd,
8894 Arc::clone(&tools),
8895 manager.clone(),
8896 resolved_policy.clone(),
8897 runtime_repair_mode,
8898 memory_limit_bytes,
8899 )
8900 .await?
8901 } else {
8902 tracing::info!(
8903 event = "pi.extension_runtime.engine_decision",
8904 stage = "agent_enable_extensions_prewarmed",
8905 requested = "native-rust",
8906 selected = "native-rust",
8907 fallback = false,
8908 "Using pre-warmed extension runtime"
8909 );
8910 ExtensionRuntimeHandle::NativeRust(runtime)
8911 }
8912 }
8913 ExtensionRuntimeHandle::Js(runtime) => {
8914 if wants_js_runtime {
8915 tracing::info!(
8916 event = "pi.extension_runtime.engine_decision",
8917 stage = "agent_enable_extensions_prewarmed",
8918 requested = "quickjs",
8919 selected = "quickjs",
8920 fallback = false,
8921 "Using pre-warmed extension runtime"
8922 );
8923 ExtensionRuntimeHandle::Js(runtime)
8924 } else {
8925 tracing::warn!(
8926 event = "pi.extension_runtime.prewarm.mismatch",
8927 expected = "native-rust",
8928 got = "quickjs",
8929 "Pre-warmed runtime mismatched requested native mode; creating native-rust runtime"
8930 );
8931 Self::start_native_extension_runtime(
8932 "agent_enable_extensions_prewarm_mismatch",
8933 cwd,
8934 Arc::clone(&tools),
8935 manager.clone(),
8936 resolved_policy.clone(),
8937 runtime_repair_mode,
8938 memory_limit_bytes,
8939 )
8940 .await?
8941 }
8942 }
8943 };
8944 manager.set_runtime(runtime);
8945 (manager, tools)
8946 } else {
8947 let manager = ExtensionManager::new();
8948 manager.set_cwd(cwd.display().to_string());
8949 let tools = Arc::new(ToolRegistry::new(enabled_tools, cwd, config));
8950
8951 if let Some(cfg) = config {
8952 let resolved_risk = cfg.resolve_extension_risk_with_metadata();
8953 tracing::info!(
8954 event = "pi.extension_runtime_risk.config",
8955 source = resolved_risk.source,
8956 enabled = resolved_risk.settings.enabled,
8957 alpha = resolved_risk.settings.alpha,
8958 window_size = resolved_risk.settings.window_size,
8959 ledger_limit = resolved_risk.settings.ledger_limit,
8960 fail_closed = resolved_risk.settings.fail_closed,
8961 "Resolved extension runtime risk settings"
8962 );
8963 manager.set_runtime_risk_config(resolved_risk.settings);
8964 }
8965
8966 let runtime = if wants_js_runtime {
8967 Self::start_js_extension_runtime(
8968 "agent_enable_extensions_boot",
8969 cwd,
8970 Arc::clone(&tools),
8971 manager.clone(),
8972 resolved_policy.clone(),
8973 runtime_repair_mode,
8974 memory_limit_bytes,
8975 )
8976 .await?
8977 } else {
8978 Self::start_native_extension_runtime(
8979 "agent_enable_extensions_boot",
8980 cwd,
8981 Arc::clone(&tools),
8982 manager.clone(),
8983 resolved_policy.clone(),
8984 runtime_repair_mode,
8985 memory_limit_bytes,
8986 )
8987 .await?
8988 };
8989 manager.set_runtime(runtime);
8990 (manager, tools)
8991 };
8992
8993 let (steering_mode, follow_up_mode) = self.agent.queue_modes();
8997 let queue_modes = Arc::new(StdMutex::new(ExtensionQueueModeState::new(
8998 steering_mode,
8999 follow_up_mode,
9000 )));
9001 manager.set_session(Arc::new(AgentExtensionSession {
9002 handle: SessionHandle(self.session.clone()),
9003 is_streaming: Arc::clone(&self.extensions_is_streaming),
9004 is_compacting: Arc::clone(&self.extensions_is_compacting),
9005 queue_modes: Arc::clone(&queue_modes),
9006 auto_compaction_enabled: self.compaction_settings.enabled,
9007 }));
9008
9009 let injected = Arc::new(StdMutex::new(ExtensionInjectedQueue::new(
9010 steering_mode,
9011 follow_up_mode,
9012 )));
9013 let host_actions = AgentSessionHostActions {
9014 session: Arc::clone(&self.session),
9015 injected: Arc::clone(&injected),
9016 is_streaming: Arc::clone(&self.extensions_is_streaming),
9017 is_turn_active: Arc::clone(&self.extensions_turn_active),
9018 pending_idle_actions: Arc::clone(&self.extensions_pending_idle_actions),
9019 ai_completion: Arc::clone(&self.extension_ai_completion),
9020 };
9021 self.extension_queue_modes = Some(Arc::clone(&queue_modes));
9022 self.extension_injected_queue = Some(Arc::clone(&injected));
9023 manager.set_host_actions(Arc::new(host_actions));
9024 {
9025 let steering_queue = Arc::clone(&injected);
9026 let follow_up_queue = Arc::clone(&injected);
9027 let steering_fetcher = move || -> BoxFuture<'static, Vec<Message>> {
9028 let steering_queue = Arc::clone(&steering_queue);
9029 Box::pin(async move {
9030 let Ok(mut queue) = steering_queue.lock() else {
9031 return Vec::new();
9032 };
9033 queue.pop_steering()
9034 })
9035 };
9036 let follow_up_fetcher = move || -> BoxFuture<'static, Vec<Message>> {
9037 let follow_up_queue = Arc::clone(&follow_up_queue);
9038 Box::pin(async move {
9039 let Ok(mut queue) = follow_up_queue.lock() else {
9040 return Vec::new();
9041 };
9042 queue.pop_follow_up()
9043 })
9044 };
9045 self.agent.register_message_fetchers(
9046 Some(Arc::new(steering_fetcher)),
9047 Some(Arc::new(follow_up_fetcher)),
9048 );
9049 }
9050
9051 if !js_specs.is_empty() {
9052 manager.load_js_extensions(js_specs).await?;
9053 }
9054
9055 if !native_specs.is_empty() {
9056 manager.load_native_extensions(native_specs).await?;
9057 }
9058
9059 if let Some(rt) = manager.runtime() {
9061 let events = rt.drain_repair_events().await;
9062 if !events.is_empty() {
9063 log_repair_diagnostics(&events);
9064 }
9065 }
9066
9067 #[cfg(feature = "wasm-host")]
9068 if !wasm_specs.is_empty() {
9069 let host = WasmExtensionHost::new(cwd, resolved_policy.clone())?;
9070 manager
9071 .load_wasm_extensions(&host, wasm_specs, Arc::clone(&tools))
9072 .await?;
9073 }
9074
9075 let session_path = {
9078 let cx = crate::agent_cx::AgentCx::for_request();
9079 let session = self
9080 .session
9081 .lock(cx.cx())
9082 .await
9083 .map_err(|e| Error::extension(e.to_string()))?;
9084 session.path.as_ref().map(|p| p.display().to_string())
9085 };
9086
9087 if let Err(err) = manager
9088 .dispatch_event(
9089 ExtensionEventName::Startup,
9090 Some(serde_json::json!({
9091 "version": env!("CARGO_PKG_VERSION"),
9092 "sessionFile": session_path,
9093 })),
9094 )
9095 .await
9096 {
9097 tracing::warn!("startup extension hook failed (fail-open): {err}");
9098 }
9099
9100 if let Err(err) = manager
9101 .dispatch_event(ExtensionEventName::SessionStart, None)
9102 .await
9103 {
9104 tracing::warn!("session_start extension hook failed (fail-open): {err}");
9105 }
9106
9107 let ctx_payload = serde_json::json!({ "cwd": cwd.display().to_string() });
9108 let wrappers = collect_extension_tool_wrappers(&manager, ctx_payload).await?;
9109 self.agent.extend_tools(wrappers);
9110 self.agent.extensions = Some(manager.clone());
9111 self.extensions = Some(ExtensionRegion::new(manager));
9112 Ok(())
9113 }
9114
9115 pub async fn save_and_index(&mut self) -> Result<()> {
9116 if self.save_enabled {
9117 let cx = crate::agent_cx::AgentCx::for_request();
9118 let mut session = OwnedMutexGuard::lock(Arc::clone(&self.session), cx.cx())
9119 .await
9120 .map_err(|e| Error::session(e.to_string()))?;
9121 session
9122 .flush_autosave(AutosaveFlushTrigger::Periodic)
9123 .await?;
9124 }
9125 Ok(())
9126 }
9127
9128 pub async fn persist_session(&mut self) -> Result<()> {
9129 if !self.save_enabled {
9130 return Ok(());
9131 }
9132 let cx = crate::agent_cx::AgentCx::for_request();
9133 let mut session = OwnedMutexGuard::lock(Arc::clone(&self.session), cx.cx())
9134 .await
9135 .map_err(|e| Error::session(e.to_string()))?;
9136 session
9137 .flush_autosave(AutosaveFlushTrigger::Periodic)
9138 .await?;
9139 Ok(())
9140 }
9141
9142 pub async fn run_text(
9143 &mut self,
9144 input: String,
9145 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9146 ) -> Result<AssistantMessage> {
9147 self.run_text_with_abort(input, None, on_event).await
9148 }
9149
9150 pub async fn run_text_with_abort(
9151 &mut self,
9152 input: String,
9153 abort: Option<AbortSignal>,
9154 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9155 ) -> Result<AssistantMessage> {
9156 self.extensions_turn_active.store(true, Ordering::SeqCst);
9157 let result = async {
9158 let outcome = self.dispatch_input_event(input, Vec::new()).await?;
9159 let (text, images) = match outcome {
9160 InputEventOutcome::Continue { text, images } => (text, images),
9161 InputEventOutcome::Block { reason } => {
9162 let message = reason.unwrap_or_else(|| "Input blocked".to_string());
9163 return Err(Error::extension(message));
9164 }
9165 };
9166
9167 let base_system_prompt = self.agent.system_prompt().map(str::to_string);
9168 let BeforeAgentStartOutcome {
9169 messages: custom_messages,
9170 system_prompt,
9171 } = self
9172 .dispatch_before_agent_start(
9173 &text,
9174 &images,
9175 base_system_prompt.as_deref().unwrap_or(""),
9176 )
9177 .await;
9178 if let Some(prompt) = system_prompt {
9179 self.agent.set_system_prompt(Some(prompt));
9180 } else {
9181 self.agent.set_system_prompt(base_system_prompt.clone());
9182 }
9183
9184 let result = if images.is_empty() {
9185 self.run_agent_with_text(text, abort, on_event, custom_messages)
9186 .await
9187 } else {
9188 let content = Self::build_content_blocks_for_input(&text, &images);
9189 self.run_agent_with_content(content, abort, on_event, custom_messages)
9190 .await
9191 };
9192
9193 self.agent.set_system_prompt(base_system_prompt);
9194 result
9195 }
9196 .await;
9197 self.extensions_turn_active.store(false, Ordering::SeqCst);
9198 result
9199 }
9200
9201 pub async fn run_with_content(
9202 &mut self,
9203 content: Vec<ContentBlock>,
9204 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9205 ) -> Result<AssistantMessage> {
9206 self.run_with_content_with_abort(content, None, on_event)
9207 .await
9208 }
9209
9210 pub async fn run_with_content_with_abort(
9211 &mut self,
9212 content: Vec<ContentBlock>,
9213 abort: Option<AbortSignal>,
9214 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9215 ) -> Result<AssistantMessage> {
9216 self.extensions_turn_active.store(true, Ordering::SeqCst);
9217 let result = async {
9218 let (text, images) = Self::split_content_blocks_for_input(&content);
9219 let outcome = self.dispatch_input_event(text, images).await?;
9220 let (text, images) = match outcome {
9221 InputEventOutcome::Continue { text, images } => (text, images),
9222 InputEventOutcome::Block { reason } => {
9223 let message = reason.unwrap_or_else(|| "Input blocked".to_string());
9224 return Err(Error::extension(message));
9225 }
9226 };
9227
9228 let base_system_prompt = self.agent.system_prompt().map(str::to_string);
9229 let BeforeAgentStartOutcome {
9230 messages: custom_messages,
9231 system_prompt,
9232 } = self
9233 .dispatch_before_agent_start(
9234 &text,
9235 &images,
9236 base_system_prompt.as_deref().unwrap_or(""),
9237 )
9238 .await;
9239 if let Some(prompt) = system_prompt {
9240 self.agent.set_system_prompt(Some(prompt));
9241 } else {
9242 self.agent.set_system_prompt(base_system_prompt.clone());
9243 }
9244
9245 let content_for_agent = Self::build_content_blocks_for_input(&text, &images);
9246 let result = self
9247 .run_agent_with_content(content_for_agent, abort, on_event, custom_messages)
9248 .await;
9249
9250 self.agent.set_system_prompt(base_system_prompt);
9251 result
9252 }
9253 .await;
9254 self.extensions_turn_active.store(false, Ordering::SeqCst);
9255 result
9256 }
9257
9258 pub async fn revert_last_user_message(&mut self) -> Result<bool> {
9259 let cx = crate::agent_cx::AgentCx::for_request();
9260 let mut session = self
9261 .session
9262 .lock(cx.cx())
9263 .await
9264 .map_err(|e| Error::session(e.to_string()))?;
9265
9266 let reverted = session.revert_last_user_message();
9267 if reverted {
9268 let messages = session.to_messages_for_current_path();
9269 self.agent.replace_messages(messages);
9270 }
9271 Ok(reverted)
9272 }
9273
9274 pub async fn revert_incomplete_response(&mut self) -> Result<bool> {
9282 let cx = crate::agent_cx::AgentCx::for_request();
9283 let mut session = self
9284 .session
9285 .lock(cx.cx())
9286 .await
9287 .map_err(|e| Error::session(e.to_string()))?;
9288
9289 let reverted = session.revert_incomplete_response();
9290 if reverted {
9291 let messages = session.to_messages_for_current_path();
9292 self.agent.replace_messages(messages);
9293 }
9294 Ok(reverted)
9295 }
9296
9297 async fn dispatch_input_event(
9298 &self,
9299 text: String,
9300 images: Vec<ImageContent>,
9301 ) -> Result<InputEventOutcome> {
9302 let Some(region) = &self.extensions else {
9303 return Ok(InputEventOutcome::Continue { text, images });
9304 };
9305
9306 let images_value = serde_json::to_value(&images).unwrap_or(Value::Null);
9307 let attachments_value = images_value.clone();
9308 let text_clone = text.clone();
9309 let payload = json!({
9310 "text": text,
9311 "content": text_clone,
9312 "images": images_value,
9313 "attachments": attachments_value,
9314 "source": self.input_source.as_str(),
9315 });
9316
9317 let response = region
9318 .manager()
9319 .dispatch_event_with_response(
9320 ExtensionEventName::Input,
9321 Some(payload),
9322 EXTENSION_EVENT_TIMEOUT_MS,
9323 )
9324 .await?;
9325
9326 Ok(apply_input_event_response(response, text, images))
9327 }
9328
9329 async fn dispatch_before_agent_start(
9330 &self,
9331 prompt: &str,
9332 images: &[ImageContent],
9333 system_prompt: &str,
9334 ) -> BeforeAgentStartOutcome {
9335 let Some(region) = &self.extensions else {
9336 return BeforeAgentStartOutcome {
9337 messages: Vec::new(),
9338 system_prompt: None,
9339 };
9340 };
9341
9342 let images_value = serde_json::to_value(images).unwrap_or(Value::Null);
9343 let payload = json!({
9344 "prompt": prompt,
9345 "images": images_value,
9346 "systemPrompt": system_prompt,
9347 });
9348
9349 let response = region
9350 .manager()
9351 .dispatch_event_with_response(
9352 ExtensionEventName::BeforeAgentStart,
9353 Some(payload),
9354 EXTENSION_EVENT_TIMEOUT_MS,
9355 )
9356 .await;
9357
9358 match response {
9359 Ok(value) => apply_before_agent_start_response(value, Utc::now().timestamp_millis()),
9360 Err(err) => {
9361 tracing::warn!("before_agent_start extension hook failed (fail-open): {err}");
9362 BeforeAgentStartOutcome {
9363 messages: Vec::new(),
9364 system_prompt: None,
9365 }
9366 }
9367 }
9368 }
9369
9370 async fn dispatch_before_compact(
9371 &self,
9372 preparation: &compaction::CompactionPreparation,
9373 branch_entries: &[crate::session::SessionEntry],
9374 custom_instructions: Option<&str>,
9375 ) -> SessionBeforeCompactOutcome {
9376 let Some(region) = &self.extensions else {
9377 return SessionBeforeCompactOutcome::default();
9378 };
9379
9380 let prep_value = compaction::compaction_preparation_to_value(preparation);
9381 let branch_entries_value =
9382 serde_json::to_value(branch_entries).unwrap_or(Value::Array(Vec::new()));
9383 let mut payload = serde_json::Map::new();
9384 payload.insert("preparation".to_string(), prep_value);
9385 payload.insert("branchEntries".to_string(), branch_entries_value);
9386 if let Some(custom_instructions) = custom_instructions {
9387 payload.insert(
9388 "customInstructions".to_string(),
9389 Value::String(custom_instructions.to_string()),
9390 );
9391 }
9392
9393 let response = region
9394 .manager()
9395 .dispatch_event_with_response(
9396 ExtensionEventName::SessionBeforeCompact,
9397 Some(Value::Object(payload)),
9398 EXTENSION_EVENT_TIMEOUT_MS,
9399 )
9400 .await;
9401
9402 match response {
9403 Ok(value) => apply_session_before_compact_response(value, preparation.tokens_before),
9404 Err(err) => {
9405 tracing::warn!("session_before_compact extension hook failed (fail-open): {err}");
9406 SessionBeforeCompactOutcome::default()
9407 }
9408 }
9409 }
9410
9411 fn prepare_semantic_context_prompt(&self) -> Option<PreparedSemanticContextPrompt> {
9412 let injection = self.semantic_context_bundle.as_ref()?;
9413 if !injection.enabled {
9414 return None;
9415 }
9416
9417 let provider = self.agent.provider();
9418 let shape = semantic_context_prompt_shape_for_provider(provider.api());
9419 let budget = semantic_context_prompt_budget_for_provider(provider.api(), injection);
9420 let revision = semantic_context_bundle_revision(&injection.bundle);
9421 let (prompt, stats) =
9422 render_semantic_context_prompt(&injection.bundle, injection, budget, &revision);
9423 if prompt.trim().is_empty() {
9424 tracing::warn!(
9425 event = "pi.semantic_context.prompt.skipped",
9426 provider = provider.name(),
9427 api = provider.api(),
9428 model = provider.model_id(),
9429 revision = %revision,
9430 max_bytes = budget.max_bytes,
9431 "semantic context bundle prompt skipped because prompt budget was too small"
9432 );
9433 return None;
9434 }
9435
9436 tracing::info!(
9437 event = "pi.semantic_context.prompt.injected",
9438 provider = provider.name(),
9439 api = provider.api(),
9440 model = provider.model_id(),
9441 revision = %revision,
9442 shape = ?shape,
9443 prompt_bytes = prompt.len(),
9444 selected_items = stats.selected_items_included,
9445 selected_items_omitted = stats.selected_items_omitted,
9446 validation_commands = stats.validation_commands_included,
9447 truncated = stats.truncated,
9448 "semantic context bundle attached to agent turn"
9449 );
9450
9451 let details = json!({
9452 "schema": SEMANTIC_CONTEXT_PROVENANCE_SCHEMA_V1,
9453 "bundleSchema": injection.bundle.schema.as_str(),
9454 "bundleRevision": revision.as_str(),
9455 "provider": {
9456 "name": provider.name(),
9457 "api": provider.api(),
9458 "model": provider.model_id(),
9459 "promptShape": shape,
9460 },
9461 "budget": {
9462 "requestedMaxItems": injection.max_prompt_items,
9463 "requestedMaxBytes": injection.max_prompt_bytes,
9464 "effectiveMaxItems": budget.max_items,
9465 "effectiveMaxBytes": budget.max_bytes,
9466 },
9467 "prompt": {
9468 "bytes": prompt.len(),
9469 "selectedItemsIncluded": stats.selected_items_included,
9470 "selectedItemsOmitted": stats.selected_items_omitted,
9471 "validationCommandsIncluded": stats.validation_commands_included,
9472 "validationCommandsOmitted": stats.validation_commands_omitted,
9473 "exclusionsIncluded": stats.exclusions_included,
9474 "exclusionsOmitted": stats.exclusions_omitted,
9475 "truncated": stats.truncated,
9476 },
9477 "bundle": {
9478 "selectedItems": injection.bundle.selected_items.len(),
9479 "excludedItems": injection.bundle.excluded_items.len(),
9480 "staleEvidenceSuppressions": injection.bundle.stale_evidence_suppressions.len(),
9481 "estimatedBytes": injection.bundle.estimated_bytes,
9482 "estimatedTokens": injection.bundle.estimated_tokens,
9483 "redactionStatus": injection.bundle.redaction_summary.overall_status,
9484 "inputFingerprintSha256": injection.bundle.invalidation_policy.input_fingerprint_sha256.as_str(),
9485 "cacheable": injection.bundle.invalidation_policy.cacheable,
9486 "workspaceId": injection.bundle.invalidation_policy.workspace_id.as_str(),
9487 "branch": injection.bundle.invalidation_policy.branch.as_deref(),
9488 "sessionId": injection.bundle.invalidation_policy.session_id.as_deref(),
9489 }
9490 });
9491
9492 Some(PreparedSemanticContextPrompt {
9493 prompt,
9494 revision,
9495 shape,
9496 details,
9497 })
9498 }
9499
9500 fn semantic_context_prompt_messages(
9501 prepared: &PreparedSemanticContextPrompt,
9502 timestamp: i64,
9503 ) -> Vec<Message> {
9504 match prepared.shape {
9505 SemanticContextPromptShape::CustomUserMessage => {
9506 vec![Message::Custom(CustomMessage {
9507 content: prepared.prompt.clone(),
9508 custom_type: SEMANTIC_CONTEXT_CUSTOM_TYPE.to_string(),
9509 display: true,
9510 details: Some(prepared.details.clone()),
9511 timestamp,
9512 })]
9513 }
9514 SemanticContextPromptShape::SystemPromptAppend => {
9515 vec![Message::Custom(CustomMessage {
9516 content: format!(
9517 "Semantic context bundle revision {} attached to system prompt.",
9518 prepared.revision
9519 ),
9520 custom_type: SEMANTIC_CONTEXT_CUSTOM_TYPE.to_string(),
9521 display: false,
9522 details: Some(prepared.details.clone()),
9523 timestamp,
9524 })]
9525 }
9526 }
9527 }
9528
9529 fn semantic_context_system_prompt_for_turn(
9530 base_system_prompt: Option<String>,
9531 prepared: Option<&PreparedSemanticContextPrompt>,
9532 ) -> Option<String> {
9533 let Some(prepared) = prepared else {
9534 return base_system_prompt;
9535 };
9536 if !matches!(
9537 prepared.shape,
9538 SemanticContextPromptShape::SystemPromptAppend
9539 ) {
9540 return base_system_prompt;
9541 }
9542
9543 let mut prompt = base_system_prompt.unwrap_or_default();
9544 if !prompt.is_empty() {
9545 prompt.push_str("\n\n");
9546 }
9547 prompt.push_str(&prepared.prompt);
9548 Some(prompt)
9549 }
9550
9551 fn split_content_blocks_for_input(blocks: &[ContentBlock]) -> (String, Vec<ImageContent>) {
9552 let mut text = String::new();
9553 let mut images = Vec::new();
9554 for block in blocks {
9555 match block {
9556 ContentBlock::Text(text_block) if !text_block.text.trim().is_empty() => {
9557 if !text.is_empty() {
9558 text.push('\n');
9559 }
9560 text.push_str(&text_block.text);
9561 }
9562 ContentBlock::Image(image) => images.push(image.clone()),
9563 _ => {}
9564 }
9565 }
9566 (text, images)
9567 }
9568
9569 fn build_content_blocks_for_input(text: &str, images: &[ImageContent]) -> Vec<ContentBlock> {
9570 let mut content = Vec::new();
9571 if !text.trim().is_empty() {
9572 content.push(ContentBlock::Text(TextContent::new(text.to_string())));
9573 }
9574 for image in images {
9575 content.push(ContentBlock::Image(image.clone()));
9576 }
9577 content
9578 }
9579
9580 fn take_pending_idle_actions(&self) -> Vec<PendingIdleAction> {
9581 let Ok(mut actions) = self.extensions_pending_idle_actions.lock() else {
9582 return Vec::new();
9583 };
9584 actions.drain(..).collect()
9585 }
9586
9587 async fn run_pending_idle_actions_with_abort(
9588 &mut self,
9589 abort: Option<AbortSignal>,
9590 on_event: AgentEventHandler,
9591 ) -> Result<()> {
9592 let actions = self.take_pending_idle_actions();
9593 if actions.is_empty() {
9594 return Ok(());
9595 }
9596
9597 let previous_source = self.input_source;
9598 self.input_source = InputSource::Extension;
9599 let result = async {
9600 for action in actions {
9601 match action {
9602 PendingIdleAction::CustomMessage(message) => {
9603 let handler = Arc::clone(&on_event);
9604 self.run_custom_message_with_abort(message, abort.clone(), move |event| {
9605 handler(event);
9606 })
9607 .await?;
9608 }
9609 PendingIdleAction::UserText(text) => {
9610 let handler = Arc::clone(&on_event);
9611 self.run_text_with_abort(text, abort.clone(), move |event| {
9612 handler(event);
9613 })
9614 .await?;
9615 }
9616 }
9617 }
9618 Ok(())
9619 }
9620 .await;
9621 self.input_source = previous_source;
9622 result
9623 }
9624
9625 async fn run_custom_message_with_abort(
9626 &mut self,
9627 message: Message,
9628 abort: Option<AbortSignal>,
9629 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9630 ) -> Result<AssistantMessage> {
9631 self.extensions_turn_active.store(true, Ordering::SeqCst);
9632 let result = async {
9633 let base_system_prompt = self.agent.system_prompt().map(str::to_string);
9634 let BeforeAgentStartOutcome {
9635 messages: custom_messages,
9636 system_prompt,
9637 } = self
9638 .dispatch_before_agent_start("", &[], base_system_prompt.as_deref().unwrap_or(""))
9639 .await;
9640 if let Some(prompt) = system_prompt {
9641 self.agent.set_system_prompt(Some(prompt));
9642 } else {
9643 self.agent.set_system_prompt(base_system_prompt.clone());
9644 }
9645
9646 let result = self
9647 .run_agent_with_prompt_message(message, abort, on_event, custom_messages)
9648 .await;
9649
9650 self.agent.set_system_prompt(base_system_prompt);
9651 result
9652 }
9653 .await;
9654 self.extensions_turn_active.store(false, Ordering::SeqCst);
9655 result
9656 }
9657
9658 async fn run_agent_with_prompt_message(
9659 &mut self,
9660 prompt_message: Message,
9661 abort: Option<AbortSignal>,
9662 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9663 custom_messages: Vec<CustomMessage>,
9664 ) -> Result<AssistantMessage> {
9665 let on_event: AgentEventHandler = Arc::new(on_event);
9666 self.sync_runtime_selection_from_session_header().await?;
9667
9668 self.maybe_compact(Arc::clone(&on_event)).await?;
9669 let history = {
9670 let cx = crate::agent_cx::AgentCx::for_request();
9671 let session = self
9672 .session
9673 .lock(cx.cx())
9674 .await
9675 .map_err(|e| Error::session(e.to_string()))?;
9676 session.to_messages_for_current_path()
9677 };
9678 self.agent.replace_messages(history);
9679
9680 let start_len = self.agent.messages().len();
9681 let mut prompts = Vec::with_capacity(1 + custom_messages.len());
9682 prompts.push(prompt_message.clone());
9683 prompts.extend(custom_messages.into_iter().map(Message::Custom));
9684
9685 {
9686 let cx = crate::agent_cx::AgentCx::for_request();
9687 let mut session = OwnedMutexGuard::lock(Arc::clone(&self.session), cx.cx())
9688 .await
9689 .map_err(|e| Error::session(e.to_string()))?;
9690 session.append_model_message(prompt_message.clone());
9691 if self.save_enabled {
9692 session.flush_autosave(AutosaveFlushTrigger::Manual).await?;
9693 }
9694 }
9695
9696 let semantic_context = self.prepare_semantic_context_prompt();
9697 let semantic_context_messages = semantic_context
9698 .as_ref()
9699 .map(|prepared| {
9700 Self::semantic_context_prompt_messages(prepared, Utc::now().timestamp_millis())
9701 })
9702 .unwrap_or_default();
9703 let streaming_guard = AtomicBoolGuard::activate(&self.extensions_is_streaming);
9704 let base_system_prompt = self.agent.system_prompt().map(str::to_string);
9705 self.agent
9706 .set_system_prompt(Self::semantic_context_system_prompt_for_turn(
9707 base_system_prompt.clone(),
9708 semantic_context.as_ref(),
9709 ));
9710 let on_event_for_run = Arc::clone(&on_event);
9711 prompts.extend(semantic_context_messages);
9712 let result = self
9713 .agent
9714 .run_with_messages_with_abort(prompts, abort, move |event| {
9715 on_event_for_run(event);
9716 })
9717 .await;
9718 drop(streaming_guard);
9719 self.agent.set_system_prompt(base_system_prompt);
9720
9721 let persist_result = self
9722 .persist_new_messages(start_len + 1, result.is_err())
9723 .await;
9724
9725 let result = result?;
9726 persist_result?;
9727 Ok(result)
9728 }
9729
9730 pub(crate) async fn run_agent_with_text(
9731 &mut self,
9732 input: String,
9733 abort: Option<AbortSignal>,
9734 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9735 custom_messages: Vec<CustomMessage>,
9736 ) -> Result<AssistantMessage> {
9737 let on_event: AgentEventHandler = Arc::new(on_event);
9738 self.sync_runtime_selection_from_session_header().await?;
9739
9740 self.maybe_compact(Arc::clone(&on_event)).await?;
9741 let history = {
9742 let cx = crate::agent_cx::AgentCx::for_request();
9743 let session = self
9744 .session
9745 .lock(cx.cx())
9746 .await
9747 .map_err(|e| Error::session(e.to_string()))?;
9748 session.to_messages_for_current_path()
9749 };
9750 self.agent.replace_messages(history);
9751
9752 let start_len = self.agent.messages().len();
9753
9754 let user_message = Message::User(UserMessage {
9756 content: UserContent::Text(input),
9757 timestamp: Utc::now().timestamp_millis(),
9758 });
9759 let mut prompts = Vec::with_capacity(1 + custom_messages.len());
9760 prompts.push(user_message.clone());
9761 let semantic_context = self.prepare_semantic_context_prompt();
9762 let semantic_context_messages = semantic_context
9763 .as_ref()
9764 .map(|prepared| {
9765 Self::semantic_context_prompt_messages(prepared, Utc::now().timestamp_millis())
9766 })
9767 .unwrap_or_default();
9768 prompts.extend(semantic_context_messages);
9769 prompts.extend(custom_messages.into_iter().map(Message::Custom));
9770
9771 {
9772 let cx = crate::agent_cx::AgentCx::for_request();
9773 let mut session = OwnedMutexGuard::lock(Arc::clone(&self.session), cx.cx())
9777 .await
9778 .map_err(|e| Error::session(e.to_string()))?;
9779 session.append_model_message(user_message.clone());
9780 if self.save_enabled {
9781 session.flush_autosave(AutosaveFlushTrigger::Manual).await?;
9782 }
9783 }
9784
9785 let streaming_guard = AtomicBoolGuard::activate(&self.extensions_is_streaming);
9786 let base_system_prompt = self.agent.system_prompt().map(str::to_string);
9787 self.agent
9788 .set_system_prompt(Self::semantic_context_system_prompt_for_turn(
9789 base_system_prompt.clone(),
9790 semantic_context.as_ref(),
9791 ));
9792 let on_event_for_run = Arc::clone(&on_event);
9793 let result = self
9794 .agent
9795 .run_with_messages_with_abort(prompts, abort, move |event| {
9796 on_event_for_run(event);
9797 })
9798 .await;
9799 drop(streaming_guard);
9800 self.agent.set_system_prompt(base_system_prompt);
9801
9802 let persist_result = self
9805 .persist_new_messages(start_len + 1, result.is_err())
9806 .await;
9807
9808 let result = result?;
9809 persist_result?;
9810 Ok(result)
9811 }
9812
9813 pub(crate) async fn run_agent_with_content(
9814 &mut self,
9815 content: Vec<ContentBlock>,
9816 abort: Option<AbortSignal>,
9817 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9818 custom_messages: Vec<CustomMessage>,
9819 ) -> Result<AssistantMessage> {
9820 let on_event: AgentEventHandler = Arc::new(on_event);
9821 self.sync_runtime_selection_from_session_header().await?;
9822
9823 self.maybe_compact(Arc::clone(&on_event)).await?;
9824 let history = {
9825 let cx = crate::agent_cx::AgentCx::for_request();
9826 let session = self
9827 .session
9828 .lock(cx.cx())
9829 .await
9830 .map_err(|e| Error::session(e.to_string()))?;
9831 session.to_messages_for_current_path()
9832 };
9833 self.agent.replace_messages(history);
9834
9835 let start_len = self.agent.messages().len();
9836
9837 let user_message = Message::User(UserMessage {
9839 content: UserContent::Blocks(content),
9840 timestamp: Utc::now().timestamp_millis(),
9841 });
9842 let mut prompts = Vec::with_capacity(1 + custom_messages.len());
9843 prompts.push(user_message.clone());
9844 let semantic_context = self.prepare_semantic_context_prompt();
9845 let semantic_context_messages = semantic_context
9846 .as_ref()
9847 .map(|prepared| {
9848 Self::semantic_context_prompt_messages(prepared, Utc::now().timestamp_millis())
9849 })
9850 .unwrap_or_default();
9851 prompts.extend(semantic_context_messages);
9852 prompts.extend(custom_messages.into_iter().map(Message::Custom));
9853
9854 {
9855 let cx = crate::agent_cx::AgentCx::for_request();
9856 let mut session = OwnedMutexGuard::lock(Arc::clone(&self.session), cx.cx())
9860 .await
9861 .map_err(|e| Error::session(e.to_string()))?;
9862 session.append_model_message(user_message.clone());
9863 if self.save_enabled {
9864 session.flush_autosave(AutosaveFlushTrigger::Manual).await?;
9865 }
9866 }
9867
9868 let streaming_guard = AtomicBoolGuard::activate(&self.extensions_is_streaming);
9869 let base_system_prompt = self.agent.system_prompt().map(str::to_string);
9870 self.agent
9871 .set_system_prompt(Self::semantic_context_system_prompt_for_turn(
9872 base_system_prompt.clone(),
9873 semantic_context.as_ref(),
9874 ));
9875 let on_event_for_run = Arc::clone(&on_event);
9876 let result = self
9877 .agent
9878 .run_with_messages_with_abort(prompts, abort, move |event| {
9879 on_event_for_run(event);
9880 })
9881 .await;
9882 drop(streaming_guard);
9883 self.agent.set_system_prompt(base_system_prompt);
9884
9885 let persist_result = self
9888 .persist_new_messages(start_len + 1, result.is_err())
9889 .await;
9890
9891 let result = result?;
9892 persist_result?;
9893 Ok(result)
9894 }
9895
9896 pub async fn run_continue_with_abort(
9908 &mut self,
9909 abort: Option<AbortSignal>,
9910 on_event: impl Fn(AgentEvent) + Send + Sync + 'static,
9911 ) -> Result<AssistantMessage> {
9912 let on_event: AgentEventHandler = Arc::new(on_event);
9913 self.sync_runtime_selection_from_session_header().await?;
9914
9915 let history = {
9918 let cx = crate::agent_cx::AgentCx::for_request();
9919 let session = self
9920 .session
9921 .lock(cx.cx())
9922 .await
9923 .map_err(|e| Error::session(e.to_string()))?;
9924 session.to_messages_for_current_path()
9925 };
9926 self.agent.replace_messages(history);
9927 let start_len = self.agent.messages().len();
9928
9929 let streaming_guard = AtomicBoolGuard::activate(&self.extensions_is_streaming);
9930 let on_event_for_run = Arc::clone(&on_event);
9931 let result = self
9932 .agent
9933 .run_continue_with_abort(abort, move |event| {
9934 on_event_for_run(event);
9935 })
9936 .await;
9937 drop(streaming_guard);
9938
9939 let persist_result = self.persist_new_messages(start_len, result.is_err()).await;
9942
9943 let result = result?;
9944 persist_result?;
9945 Ok(result)
9946 }
9947
9948 async fn persist_new_messages(&self, start_len: usize, run_failed: bool) -> Result<()> {
9949 let new_messages = self.agent.messages()[start_len..].to_vec();
9950 {
9951 let cx = crate::agent_cx::AgentCx::for_request();
9952 let mut session = OwnedMutexGuard::lock(Arc::clone(&self.session), cx.cx())
9953 .await
9954 .map_err(|e| Error::session(e.to_string()))?;
9955 for message in new_messages {
9956 if run_failed && is_synthetic_empty_error_assistant(&message) {
9957 continue;
9958 }
9959 session.append_model_message(message);
9960 }
9961 if self.save_enabled {
9962 session
9963 .flush_autosave(AutosaveFlushTrigger::Periodic)
9964 .await?;
9965 }
9966 }
9967 Ok(())
9968 }
9969}
9970
9971fn is_synthetic_empty_error_assistant(message: &Message) -> bool {
9972 matches!(
9973 message,
9974 Message::Assistant(assistant)
9975 if assistant.content.is_empty()
9976 && matches!(assistant.stop_reason, StopReason::Error)
9977 && assistant.error_message.is_some()
9978 )
9979}
9980
9981fn semantic_context_prompt_shape_for_provider(api: &str) -> SemanticContextPromptShape {
9982 match api {
9983 "bedrock-converse-stream" | "gitlab-chat" => SemanticContextPromptShape::SystemPromptAppend,
9984 _ => SemanticContextPromptShape::CustomUserMessage,
9985 }
9986}
9987
9988fn semantic_context_prompt_budget_for_provider(
9989 api: &str,
9990 injection: &SemanticContextBundleInjection,
9991) -> SemanticContextPromptBudget {
9992 let provider_max_bytes = match api {
9993 "gitlab-chat" => 8 * 1024,
9994 "bedrock-converse-stream" | "google-gemini" | "google-vertex" => 12 * 1024,
9995 "openai-responses" | "openai-completions" | "azure-openai" => 24 * 1024,
9996 "anthropic" => 32 * 1024,
9997 _ => DEFAULT_SEMANTIC_CONTEXT_PROMPT_MAX_BYTES,
9998 };
9999 let provider_max_items = match api {
10000 "gitlab-chat" => 8,
10001 "bedrock-converse-stream" | "google-gemini" | "google-vertex" => 12,
10002 _ => DEFAULT_SEMANTIC_CONTEXT_PROMPT_MAX_ITEMS,
10003 };
10004
10005 SemanticContextPromptBudget {
10006 max_items: injection
10007 .max_prompt_items
10008 .min(injection.bundle.budget.max_items)
10009 .min(provider_max_items),
10010 max_bytes: injection
10011 .max_prompt_bytes
10012 .min(injection.bundle.budget.max_bytes)
10013 .min(provider_max_bytes),
10014 }
10015}
10016
10017fn semantic_context_bundle_revision(bundle: &SemanticContextBundle) -> String {
10018 let bytes = serde_json::to_vec(bundle).unwrap_or_else(|_| {
10019 format!(
10020 "{}:{}:{}:{}",
10021 bundle.schema,
10022 bundle.invalidation_policy.input_fingerprint_sha256,
10023 bundle.selected_items.len(),
10024 bundle.estimated_bytes
10025 )
10026 .into_bytes()
10027 });
10028 format!("{:x}", Sha256::digest(bytes))
10029}
10030
10031fn render_semantic_context_prompt(
10032 bundle: &SemanticContextBundle,
10033 injection: &SemanticContextBundleInjection,
10034 budget: SemanticContextPromptBudget,
10035 revision: &str,
10036) -> (String, SemanticContextPromptStats) {
10037 let mut prompt = String::new();
10038 let mut stats = SemanticContextPromptStats::default();
10039 push_semantic_context_header(&mut prompt, &mut stats, budget, bundle, revision);
10040 push_selected_semantic_context_items(&mut prompt, &mut stats, budget, bundle);
10041 if injection.include_validation_commands {
10042 push_semantic_context_validation_commands(&mut prompt, &mut stats, budget, bundle);
10043 }
10044 if injection.include_exclusion_summary {
10045 push_semantic_context_exclusions(&mut prompt, &mut stats, budget, bundle);
10046 }
10047
10048 if prompt.len() > usize::try_from(budget.max_bytes).unwrap_or(usize::MAX) {
10049 stats.truncated = true;
10050 truncate_string_to_max_bytes(&mut prompt, budget.max_bytes);
10051 }
10052
10053 (prompt, stats)
10054}
10055
10056fn push_semantic_context_header(
10057 prompt: &mut String,
10058 stats: &mut SemanticContextPromptStats,
10059 budget: SemanticContextPromptBudget,
10060 bundle: &SemanticContextBundle,
10061 revision: &str,
10062) {
10063 let branch = bundle
10064 .invalidation_policy
10065 .branch
10066 .as_deref()
10067 .map_or_else(|| "(none)".to_string(), safe_context_field);
10068 let session = bundle
10069 .invalidation_policy
10070 .session_id
10071 .as_deref()
10072 .map_or_else(|| "(none)".to_string(), safe_context_field);
10073
10074 let header = format!(
10075 "# Semantic Context Bundle\nschema: {SEMANTIC_CONTEXT_PROMPT_SCHEMA_V1}\nrevision: {revision}"
10076 );
10077 push_semantic_context_line(prompt, budget.max_bytes, &header, stats);
10078 push_semantic_context_line(
10079 prompt,
10080 budget.max_bytes,
10081 "Use this as navigation context for the current turn. Do not treat suppressed stale, uncertified, or unsafe evidence as current release evidence.",
10082 stats,
10083 );
10084 push_semantic_context_line(
10085 prompt,
10086 budget.max_bytes,
10087 &format!(
10088 "bundle: schema={} estimated_bytes={} estimated_tokens={} redaction={:?}",
10089 safe_context_field(&bundle.schema),
10090 bundle.estimated_bytes,
10091 bundle.estimated_tokens,
10092 bundle.redaction_summary.overall_status
10093 ),
10094 stats,
10095 );
10096 push_semantic_context_line(
10097 prompt,
10098 budget.max_bytes,
10099 &format!(
10100 "provenance: workspace={} branch={} session={} input_fingerprint_sha256={}",
10101 safe_context_field(&bundle.invalidation_policy.workspace_id),
10102 branch,
10103 session,
10104 safe_context_field(&bundle.invalidation_policy.input_fingerprint_sha256)
10105 ),
10106 stats,
10107 );
10108}
10109
10110fn push_selected_semantic_context_items(
10111 prompt: &mut String,
10112 stats: &mut SemanticContextPromptStats,
10113 budget: SemanticContextPromptBudget,
10114 bundle: &SemanticContextBundle,
10115) {
10116 push_semantic_context_line(prompt, budget.max_bytes, "", stats);
10117 push_semantic_context_line(prompt, budget.max_bytes, "Selected context:", stats);
10118 for (index, item) in bundle.selected_items.iter().enumerate() {
10119 if index >= budget.max_items {
10120 stats.selected_items_omitted = stats
10121 .selected_items_omitted
10122 .saturating_add(bundle.selected_items.len().saturating_sub(index));
10123 break;
10124 }
10125 if push_semantic_context_item(prompt, stats, budget, item, index + 1) {
10126 stats.selected_items_included = stats.selected_items_included.saturating_add(1);
10127 } else {
10128 stats.selected_items_omitted = stats
10129 .selected_items_omitted
10130 .saturating_add(bundle.selected_items.len().saturating_sub(index));
10131 break;
10132 }
10133 }
10134 if bundle.selected_items.is_empty() {
10135 push_semantic_context_line(prompt, budget.max_bytes, "- (none)", stats);
10136 }
10137}
10138
10139fn push_semantic_context_validation_commands(
10140 prompt: &mut String,
10141 stats: &mut SemanticContextPromptStats,
10142 budget: SemanticContextPromptBudget,
10143 bundle: &SemanticContextBundle,
10144) {
10145 push_semantic_context_line(prompt, budget.max_bytes, "", stats);
10146 push_semantic_context_line(
10147 prompt,
10148 budget.max_bytes,
10149 "Suggested validation commands:",
10150 stats,
10151 );
10152 if bundle.suggested_validation_commands.is_empty() {
10153 push_semantic_context_line(prompt, budget.max_bytes, "- (none)", stats);
10154 return;
10155 }
10156
10157 for (index, command) in bundle.suggested_validation_commands.iter().enumerate() {
10158 let line = format!("- {}", safe_context_field(command));
10159 if push_semantic_context_line(prompt, budget.max_bytes, &line, stats) {
10160 stats.validation_commands_included =
10161 stats.validation_commands_included.saturating_add(1);
10162 } else {
10163 stats.validation_commands_omitted = bundle
10164 .suggested_validation_commands
10165 .len()
10166 .saturating_sub(index);
10167 break;
10168 }
10169 }
10170}
10171
10172fn push_semantic_context_exclusions(
10173 prompt: &mut String,
10174 stats: &mut SemanticContextPromptStats,
10175 budget: SemanticContextPromptBudget,
10176 bundle: &SemanticContextBundle,
10177) {
10178 push_semantic_context_line(prompt, budget.max_bytes, "", stats);
10179 push_semantic_context_line(
10180 prompt,
10181 budget.max_bytes,
10182 "Suppressed or excluded context:",
10183 stats,
10184 );
10185 if bundle.stale_evidence_suppressions.is_empty() && bundle.excluded_items.is_empty() {
10186 push_semantic_context_line(prompt, budget.max_bytes, "- (none)", stats);
10187 return;
10188 }
10189
10190 for (index, item) in bundle
10191 .stale_evidence_suppressions
10192 .iter()
10193 .chain(bundle.excluded_items.iter())
10194 .take(8)
10195 .enumerate()
10196 {
10197 let line = format!(
10198 "- {:?} {} :: {} reason={}",
10199 item.node_type,
10200 safe_context_field(&item.source_path),
10201 safe_context_field(&item.title),
10202 safe_context_field(&item.reason)
10203 );
10204 if push_semantic_context_line(prompt, budget.max_bytes, &line, stats) {
10205 stats.exclusions_included = stats.exclusions_included.saturating_add(1);
10206 } else {
10207 stats.exclusions_omitted = bundle
10208 .stale_evidence_suppressions
10209 .len()
10210 .saturating_add(bundle.excluded_items.len())
10211 .saturating_sub(index);
10212 break;
10213 }
10214 }
10215}
10216
10217fn push_semantic_context_item(
10218 prompt: &mut String,
10219 stats: &mut SemanticContextPromptStats,
10220 budget: SemanticContextPromptBudget,
10221 item: &ContextBundleItem,
10222 ordinal: usize,
10223) -> bool {
10224 let freshness = item.freshness_status.map_or_else(
10225 || "not_applicable".to_string(),
10226 |status| format!("{status:?}"),
10227 );
10228 let line = format!(
10229 "{ordinal}. {:?} {} :: {}",
10230 item.node_type,
10231 safe_context_field(&item.source_path),
10232 safe_context_field(&item.title)
10233 );
10234 let detail = format!(
10235 " reason={} score={} tokens={} freshness={} redaction={:?}",
10236 safe_context_field(&item.reason),
10237 item.score,
10238 item.estimated_tokens,
10239 freshness,
10240 item.redaction_status
10241 );
10242 push_semantic_context_line(prompt, budget.max_bytes, &line, stats)
10243 && push_semantic_context_line(prompt, budget.max_bytes, &detail, stats)
10244}
10245
10246fn push_semantic_context_line(
10247 prompt: &mut String,
10248 max_bytes: u64,
10249 line: &str,
10250 stats: &mut SemanticContextPromptStats,
10251) -> bool {
10252 let max_bytes = usize::try_from(max_bytes).unwrap_or(usize::MAX);
10253 let required = line.len().saturating_add(1);
10254 if prompt.len().saturating_add(required) > max_bytes {
10255 stats.truncated = true;
10256 return false;
10257 }
10258 prompt.push_str(line);
10259 prompt.push('\n');
10260 true
10261}
10262
10263fn truncate_string_to_max_bytes(value: &mut String, max_bytes: u64) {
10264 let max_bytes = usize::try_from(max_bytes).unwrap_or(usize::MAX);
10265 if value.len() <= max_bytes {
10266 return;
10267 }
10268 let mut end = max_bytes;
10269 while !value.is_char_boundary(end) {
10270 end = end.saturating_sub(1);
10271 }
10272 value.truncate(end);
10273}
10274
10275fn safe_context_field(value: &str) -> String {
10276 let mut output = String::with_capacity(value.len().min(512));
10277 for ch in value.chars() {
10278 if matches!(ch, '\n' | '\r' | '\t') {
10279 output.push(' ');
10280 } else if ch.is_control() {
10281 output.push('?');
10282 } else {
10283 output.push(ch);
10284 }
10285 if output.len() >= 512 {
10286 output.push_str("...");
10287 break;
10288 }
10289 }
10290 output
10291}
10292
10293fn log_repair_diagnostics(events: &[crate::extensions_js::ExtensionRepairEvent]) {
10302 use std::collections::BTreeMap;
10303
10304 for ev in events {
10306 tracing::info!(
10307 event = "extension.auto_repair",
10308 extension_id = %ev.extension_id,
10309 pattern = %ev.pattern,
10310 success = ev.success,
10311 original_error = %ev.original_error,
10312 repair_action = %ev.repair_action,
10313 );
10314 }
10315
10316 let mut by_pattern: BTreeMap<String, Vec<&str>> = BTreeMap::new();
10318 for ev in events {
10319 by_pattern
10320 .entry(ev.pattern.to_string())
10321 .or_default()
10322 .push(&ev.extension_id);
10323 }
10324
10325 let verbose = std::env::var("PI_AUTO_REPAIR_VERBOSE")
10326 .is_ok_and(|v| v == "1" || v.eq_ignore_ascii_case("true"));
10327
10328 if verbose {
10329 warn!(
10330 "[auto-repair] {} extension{} auto-repaired:",
10331 events.len(),
10332 if events.len() == 1 { "" } else { "s" }
10333 );
10334 for ev in events {
10335 warn!(
10336 " {}: {} ({})",
10337 ev.pattern, ev.extension_id, ev.repair_action
10338 );
10339 }
10340 } else {
10341 let patterns: Vec<String> = by_pattern
10343 .iter()
10344 .map(|(pat, ids)| format!("{pat}:{}", ids.len()))
10345 .collect();
10346 tracing::info!(
10347 event = "extension.auto_repair.summary",
10348 count = events.len(),
10349 patterns = %patterns.join(", "),
10350 "auto-repaired {} extension(s)",
10351 events.len(),
10352 );
10353 }
10354}
10355
10356const BLOCK_IMAGES_PLACEHOLDER: &str = "Image reading is disabled.";
10357
10358#[derive(Debug, Default, Clone, Copy)]
10359struct ImageFilterStats {
10360 removed_images: usize,
10361 affected_messages: usize,
10362}
10363
10364fn filter_images_for_provider(messages: &mut [Message]) -> ImageFilterStats {
10365 let mut stats = ImageFilterStats::default();
10366 for message in messages {
10367 let removed = filter_images_from_message(message);
10368 if removed > 0 {
10369 stats.removed_images += removed;
10370 stats.affected_messages += 1;
10371 }
10372 }
10373 stats
10374}
10375
10376fn filter_images_from_message(message: &mut Message) -> usize {
10377 match message {
10378 Message::User(user) => match &mut user.content {
10379 UserContent::Text(_) => 0,
10380 UserContent::Blocks(blocks) => filter_image_blocks(blocks),
10381 },
10382 Message::Assistant(assistant) => {
10383 let assistant = Arc::make_mut(assistant);
10384 filter_image_blocks(&mut assistant.content)
10385 }
10386 Message::ToolResult(tool_result) => {
10387 filter_image_blocks(&mut Arc::make_mut(tool_result).content)
10388 }
10389 Message::Custom(_) => 0,
10390 }
10391}
10392
10393fn filter_image_blocks(blocks: &mut Vec<ContentBlock>) -> usize {
10394 let mut removed = 0usize;
10395 let mut filtered = Vec::with_capacity(blocks.len());
10396
10397 for block in blocks.drain(..) {
10398 match block {
10399 ContentBlock::Image(_) => {
10400 removed += 1;
10401 let previous_is_placeholder =
10402 filtered
10403 .last()
10404 .is_some_and(|prev| matches!(prev, ContentBlock::Text(TextContent { text, .. }) if text.as_str().eq(BLOCK_IMAGES_PLACEHOLDER)));
10405 if !previous_is_placeholder {
10406 filtered.push(ContentBlock::Text(TextContent::new(
10407 BLOCK_IMAGES_PLACEHOLDER,
10408 )));
10409 }
10410 }
10411 other => filtered.push(other),
10412 }
10413 }
10414
10415 *blocks = filtered;
10416 removed
10417}
10418
10419fn extract_tool_calls(content: &[ContentBlock]) -> Vec<ToolCall> {
10421 content
10422 .iter()
10423 .filter_map(|block| {
10424 if let ContentBlock::ToolCall(tc) = block {
10425 Some(tc.clone())
10426 } else {
10427 None
10428 }
10429 })
10430 .collect()
10431}
10432
10433#[cfg(test)]
10438mod tests {
10439 use super::*;
10440 use crate::auth::AuthCredential;
10441 use crate::provider::{InputType, Model, ModelCost};
10442 use asupersync::runtime::RuntimeBuilder;
10443 use async_trait::async_trait;
10444 use futures::Stream;
10445 use std::collections::BTreeSet;
10446 use std::collections::HashMap;
10447 use std::path::Path;
10448 use std::pin::Pin;
10449 use std::sync::{Arc as StdArc, Mutex as StdTestMutex};
10450
10451 fn user_message(text: &str) -> Message {
10452 Message::User(UserMessage {
10453 content: UserContent::Text(text.to_string()),
10454 timestamp: 0,
10455 })
10456 }
10457
10458 fn assert_user_text(message: &Message, expected: &str) {
10459 assert!(
10460 matches!(
10461 message,
10462 Message::User(UserMessage {
10463 content: UserContent::Text(_),
10464 ..
10465 })
10466 ),
10467 "expected user text message, got {message:?}"
10468 );
10469 if let Message::User(UserMessage {
10470 content: UserContent::Text(text),
10471 ..
10472 }) = message
10473 {
10474 assert_eq!(text, expected);
10475 }
10476 }
10477
10478 fn sample_image_block() -> ContentBlock {
10479 ContentBlock::Image(ImageContent {
10480 data: "aGVsbG8=".to_string(),
10481 mime_type: "image/png".to_string(),
10482 })
10483 }
10484
10485 fn image_count_in_message(message: &Message) -> usize {
10486 let count_images = |blocks: &[ContentBlock]| {
10487 blocks
10488 .iter()
10489 .filter(|block| matches!(block, ContentBlock::Image(_)))
10490 .count()
10491 };
10492 match message {
10493 Message::User(UserMessage {
10494 content: UserContent::Blocks(blocks),
10495 ..
10496 }) => count_images(blocks),
10497 Message::Assistant(msg) => count_images(&msg.content),
10498 Message::ToolResult(tool_result) => count_images(&tool_result.content),
10499 Message::User(UserMessage {
10500 content: UserContent::Text(_),
10501 ..
10502 })
10503 | Message::Custom(_) => 0,
10504 }
10505 }
10506
10507 fn assistant_message(text: &str) -> AssistantMessage {
10508 AssistantMessage {
10509 content: vec![ContentBlock::Text(TextContent::new(text))],
10510 api: "test-api".to_string(),
10511 provider: "test-provider".to_string(),
10512 model: "test-model".to_string(),
10513 usage: Usage::default(),
10514 stop_reason: StopReason::Stop,
10515 error_message: None,
10516 timestamp: 0,
10517 }
10518 }
10519
10520 #[derive(Debug)]
10521 struct SilentProvider;
10522
10523 #[async_trait]
10524 #[allow(clippy::unnecessary_literal_bound)]
10525 impl Provider for SilentProvider {
10526 fn name(&self) -> &str {
10527 "silent-provider"
10528 }
10529
10530 fn api(&self) -> &str {
10531 "test-api"
10532 }
10533
10534 fn model_id(&self) -> &str {
10535 "test-model"
10536 }
10537
10538 async fn stream(
10539 &self,
10540 _context: &Context<'_>,
10541 _options: &StreamOptions,
10542 ) -> crate::error::Result<
10543 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
10544 > {
10545 Ok(Box::pin(futures::stream::empty()))
10546 }
10547 }
10548
10549 #[derive(Debug)]
10550 struct DeltaOnlyProvider;
10551
10552 #[async_trait]
10553 #[allow(clippy::unnecessary_literal_bound)]
10554 impl Provider for DeltaOnlyProvider {
10555 fn name(&self) -> &str {
10556 "test-provider"
10557 }
10558
10559 fn api(&self) -> &str {
10560 "test-api"
10561 }
10562
10563 fn model_id(&self) -> &str {
10564 "test-model"
10565 }
10566
10567 async fn stream(
10568 &self,
10569 _context: &Context<'_>,
10570 _options: &StreamOptions,
10571 ) -> crate::error::Result<
10572 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
10573 > {
10574 let final_message = assistant_message("hello");
10575 let events = vec![
10576 Ok(StreamEvent::TextDelta {
10577 content_index: 0,
10578 delta: "hello".to_string(),
10579 }),
10580 Ok(StreamEvent::Done {
10581 reason: StopReason::Stop,
10582 message: final_message,
10583 }),
10584 ];
10585 Ok(Box::pin(futures::stream::iter(events)))
10586 }
10587 }
10588
10589 #[derive(Debug, Default)]
10590 struct CapturedProviderContext {
10591 system_prompt: Option<String>,
10592 messages: Vec<Message>,
10593 }
10594
10595 #[derive(Debug)]
10596 struct CapturingProvider {
10597 api: &'static str,
10598 calls: StdArc<StdTestMutex<Vec<CapturedProviderContext>>>,
10599 }
10600
10601 impl CapturingProvider {
10602 fn new(api: &'static str) -> Self {
10603 Self {
10604 api,
10605 calls: StdArc::new(StdTestMutex::new(Vec::new())),
10606 }
10607 }
10608
10609 fn calls(&self) -> StdArc<StdTestMutex<Vec<CapturedProviderContext>>> {
10610 StdArc::clone(&self.calls)
10611 }
10612 }
10613
10614 #[async_trait]
10615 #[allow(clippy::unnecessary_literal_bound)]
10616 impl Provider for CapturingProvider {
10617 fn name(&self) -> &str {
10618 "capturing-provider"
10619 }
10620
10621 fn api(&self) -> &str {
10622 self.api
10623 }
10624
10625 fn model_id(&self) -> &str {
10626 "capture-model"
10627 }
10628
10629 async fn stream(
10630 &self,
10631 context: &Context<'_>,
10632 _options: &StreamOptions,
10633 ) -> crate::error::Result<
10634 Pin<Box<dyn Stream<Item = crate::error::Result<StreamEvent>> + Send>>,
10635 > {
10636 self.calls
10637 .lock()
10638 .expect("capture context lock")
10639 .push(CapturedProviderContext {
10640 system_prompt: context.system_prompt.as_ref().map(ToString::to_string),
10641 messages: context.messages.iter().cloned().collect(),
10642 });
10643 let final_message = assistant_message("captured");
10644 Ok(Box::pin(futures::stream::iter(vec![Ok(
10645 StreamEvent::Done {
10646 reason: StopReason::Stop,
10647 message: final_message,
10648 },
10649 )])))
10650 }
10651 }
10652
10653 fn sample_semantic_context_bundle() -> SemanticContextBundle {
10654 use crate::semantic_workspace_graph::{
10655 ContextBundleBudget, ContextBundleExclusion, ContextBundleInvalidationPolicy,
10656 ContextRedactionSummary, EvidenceFreshnessStatus, RedactionStatus, SemanticNodeType,
10657 };
10658
10659 SemanticContextBundle {
10660 schema: crate::semantic_workspace_graph::SEMANTIC_CONTEXT_BUNDLE_SCHEMA.to_string(),
10661 budget: ContextBundleBudget {
10662 max_items: 8,
10663 max_bytes: 4096,
10664 },
10665 selected_items: vec![
10666 ContextBundleItem {
10667 node_id: "node-session".to_string(),
10668 node_type: SemanticNodeType::CodeSymbol,
10669 source_path: "src/agent.rs".to_string(),
10670 title: "AgentSession::run_agent_with_text".to_string(),
10671 reason: "query_match,related_to_bead_or_changed_path".to_string(),
10672 score: 420,
10673 estimated_bytes: 700,
10674 estimated_tokens: 175,
10675 freshness_status: None,
10676 redaction_status: RedactionStatus::None,
10677 },
10678 ContextBundleItem {
10679 node_id: "node-test".to_string(),
10680 node_type: SemanticNodeType::TestCase,
10681 source_path: "tests/agent_loop_reliability.rs".to_string(),
10682 title: "semantic context session coverage".to_string(),
10683 reason: "validation_context".to_string(),
10684 score: 300,
10685 estimated_bytes: 400,
10686 estimated_tokens: 100,
10687 freshness_status: Some(EvidenceFreshnessStatus::Current),
10688 redaction_status: RedactionStatus::Redacted,
10689 },
10690 ],
10691 excluded_items: vec![ContextBundleExclusion {
10692 node_id: "stale-doc".to_string(),
10693 node_type: SemanticNodeType::DocSection,
10694 source_path: "README.md".to_string(),
10695 title: "obsolete drop-in claim".to_string(),
10696 reason: "suppressed_stale_or_unsafe_evidence".to_string(),
10697 score: 250,
10698 estimated_bytes: 300,
10699 freshness_status: Some(EvidenceFreshnessStatus::Uncertified),
10700 redaction_status: RedactionStatus::SensitiveOmitted,
10701 }],
10702 stale_evidence_suppressions: Vec::new(),
10703 redaction_summary: ContextRedactionSummary {
10704 policy_version: "test-policy".to_string(),
10705 overall_status: RedactionStatus::Redacted,
10706 selected_redacted_nodes: 1,
10707 selected_sensitive_omissions: 0,
10708 suppressed_unsafe_nodes: 0,
10709 redacted_metadata_keys: BTreeSet::from(["api_key".to_string()]),
10710 sensitive_path_kinds: BTreeSet::new(),
10711 },
10712 invalidation_policy: ContextBundleInvalidationPolicy {
10713 policy_version: "test-policy".to_string(),
10714 workspace_id: "workspace:test".to_string(),
10715 branch: Some("main".to_string()),
10716 session_id: Some("session-123".to_string()),
10717 input_fingerprint_sha256: "abc123".repeat(10),
10718 cache_ttl_seconds: 900,
10719 generated_at_utc: Some("2026-05-13T00:00:00Z".to_string()),
10720 expires_at_utc: Some("2026-05-13T00:15:00Z".to_string()),
10721 invalidates_on: vec!["input_fingerprint_change".to_string()],
10722 cacheable: true,
10723 },
10724 path_normalization: Vec::new(),
10725 suggested_validation_commands: vec![
10726 "cargo test agent_semantic_context".to_string(),
10727 "cargo check --all-targets".to_string(),
10728 ],
10729 estimated_bytes: 1100,
10730 estimated_tokens: 275,
10731 }
10732 }
10733
10734 #[test]
10735 fn delta_without_start_does_not_mutate_previous_message() {
10736 let runtime = RuntimeBuilder::current_thread()
10737 .build()
10738 .expect("runtime build");
10739
10740 runtime.block_on(async {
10741 let provider = Arc::new(DeltaOnlyProvider);
10742 let tools = ToolRegistry::from_tools(Vec::new());
10743 let mut agent = Agent::new(provider, tools, AgentConfig::default());
10744
10745 agent.add_message(Message::Assistant(Arc::new(assistant_message("prev"))));
10746
10747 agent
10748 .run_with_message_with_abort(user_message("hi"), None, |_| {})
10749 .await
10750 .expect("run");
10751
10752 let assistant_texts = agent
10753 .messages()
10754 .iter()
10755 .filter_map(|message| match message {
10756 Message::Assistant(msg)
10757 if matches!(msg.content.as_slice(), [ContentBlock::Text(_)]) =>
10758 {
10759 if let [ContentBlock::Text(text)] = msg.content.as_slice() {
10760 Some(text.text.clone())
10761 } else {
10762 None
10763 }
10764 }
10765 _ => None,
10766 })
10767 .collect::<Vec<_>>();
10768
10769 assert_eq!(
10770 assistant_texts.as_slice(),
10771 ["prev".to_string(), "hello".to_string()]
10772 );
10773 });
10774 }
10775
10776 #[test]
10777 fn semantic_context_bundle_injection_is_disabled_by_default() {
10778 let runtime = RuntimeBuilder::current_thread()
10779 .build()
10780 .expect("runtime build");
10781
10782 runtime.block_on(async {
10783 let provider = CapturingProvider::new("openai-responses");
10784 let calls = provider.calls();
10785 let agent = Agent::new(
10786 Arc::new(provider),
10787 ToolRegistry::from_tools(Vec::new()),
10788 AgentConfig::default(),
10789 );
10790 let session = Arc::new(Mutex::new(Session::in_memory()));
10791 let mut agent_session =
10792 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
10793
10794 agent_session
10795 .run_text("hello".to_string(), |_| {})
10796 .await
10797 .expect("run with default context settings");
10798
10799 let calls = match calls.lock() {
10800 Ok(calls) => calls,
10801 Err(poisoned) => poisoned.into_inner(),
10802 };
10803 assert_eq!(calls.len(), 1);
10804 assert_eq!(calls[0].messages.len(), 1);
10805 assert_user_text(&calls[0].messages[0], "hello");
10806 assert!(calls[0].system_prompt.is_none());
10807 drop(calls);
10808 });
10809 }
10810
10811 #[test]
10812 fn semantic_context_bundle_injection_adds_bounded_custom_message_and_session_provenance() {
10813 let runtime = RuntimeBuilder::current_thread()
10814 .build()
10815 .expect("runtime build");
10816
10817 runtime.block_on(async {
10818 let bundle = sample_semantic_context_bundle();
10819 let revision = semantic_context_bundle_revision(&bundle);
10820 let provider = CapturingProvider::new("openai-responses");
10821 let calls = provider.calls();
10822 let agent = Agent::new(
10823 Arc::new(provider),
10824 ToolRegistry::from_tools(Vec::new()),
10825 AgentConfig::default(),
10826 );
10827 let session = Arc::new(Mutex::new(Session::in_memory()));
10828 let mut agent_session = AgentSession::new(
10829 agent,
10830 Arc::clone(&session),
10831 false,
10832 ResolvedCompactionSettings::default(),
10833 );
10834 agent_session.set_semantic_context_bundle(Some(
10835 SemanticContextBundleInjection::enabled(bundle).with_prompt_budget(4, 2048),
10836 ));
10837
10838 agent_session
10839 .run_text("use context".to_string(), |_| {})
10840 .await
10841 .expect("run with context bundle");
10842
10843 {
10844 let calls = match calls.lock() {
10845 Ok(calls) => calls,
10846 Err(poisoned) => poisoned.into_inner(),
10847 };
10848 assert_eq!(calls.len(), 1);
10849 assert_eq!(calls[0].messages.len(), 2);
10850 assert_user_text(&calls[0].messages[0], "use context");
10851 let custom = match &calls[0].messages[1] {
10852 Message::Custom(custom) => custom,
10853 other => {
10854 assert!(
10855 matches!(other, Message::Custom(_)),
10856 "expected custom semantic context message"
10857 );
10858 return;
10859 }
10860 };
10861 assert_eq!(custom.custom_type, SEMANTIC_CONTEXT_CUSTOM_TYPE);
10862 assert!(custom.display);
10863 assert!(custom.content.len() <= 2048);
10864 assert!(custom.content.contains("Semantic Context Bundle"));
10865 assert!(custom.content.contains("src/agent.rs"));
10866 let details = custom.details.as_ref().expect("context provenance");
10867 assert_eq!(
10868 details.get("bundleRevision").and_then(Value::as_str),
10869 Some(revision.as_str())
10870 );
10871 assert_eq!(
10872 details
10873 .pointer("/provider/promptShape")
10874 .and_then(Value::as_str),
10875 Some("custom_user_message")
10876 );
10877 drop(calls);
10878 }
10879
10880 let cx = crate::agent_cx::AgentCx::for_request();
10881 let stored = session
10882 .lock(cx.cx())
10883 .await
10884 .expect("session lock")
10885 .to_messages_for_current_path();
10886 assert!(
10887 stored.iter().any(|message| matches!(
10888 message,
10889 Message::Custom(CustomMessage { custom_type, details, display: true, .. })
10890 if custom_type == SEMANTIC_CONTEXT_CUSTOM_TYPE
10891 && details
10892 .as_ref()
10893 .and_then(|value| value.get("bundleRevision"))
10894 .and_then(Value::as_str)
10895 == Some(revision.as_str())
10896 )),
10897 "semantic context provenance was not persisted in session messages: {stored:?}"
10898 );
10899 });
10900 }
10901
10902 #[test]
10903 fn semantic_context_bundle_uses_system_prompt_append_for_providers_without_custom_context() {
10904 let runtime = RuntimeBuilder::current_thread()
10905 .build()
10906 .expect("runtime build");
10907
10908 runtime.block_on(async {
10909 let bundle = sample_semantic_context_bundle();
10910 let revision = semantic_context_bundle_revision(&bundle);
10911 let provider = CapturingProvider::new("gitlab-chat");
10912 let calls = provider.calls();
10913 let agent = Agent::new(
10914 Arc::new(provider),
10915 ToolRegistry::from_tools(Vec::new()),
10916 AgentConfig {
10917 system_prompt: Some("base prompt".to_string()),
10918 ..AgentConfig::default()
10919 },
10920 );
10921 let session = Arc::new(Mutex::new(Session::in_memory()));
10922 let mut agent_session = AgentSession::new(
10923 agent,
10924 Arc::clone(&session),
10925 false,
10926 ResolvedCompactionSettings::default(),
10927 );
10928 agent_session.set_semantic_context_bundle(Some(
10929 SemanticContextBundleInjection::enabled(bundle).with_prompt_budget(4, 2048),
10930 ));
10931
10932 agent_session
10933 .run_text("gitlab turn".to_string(), |_| {})
10934 .await
10935 .expect("run with system prompt context");
10936
10937 {
10938 let calls = match calls.lock() {
10939 Ok(calls) => calls,
10940 Err(poisoned) => poisoned.into_inner(),
10941 };
10942 assert_eq!(calls.len(), 1);
10943 assert_eq!(calls[0].messages.len(), 1);
10944 assert_user_text(&calls[0].messages[0], "gitlab turn");
10945 let system_prompt = calls[0].system_prompt.as_deref().expect("system prompt");
10946 assert!(system_prompt.contains("base prompt"));
10947 assert!(system_prompt.contains("Semantic Context Bundle"));
10948 assert!(system_prompt.contains("src/agent.rs"));
10949 drop(calls);
10950 }
10951
10952 let cx = crate::agent_cx::AgentCx::for_request();
10953 let stored = session
10954 .lock(cx.cx())
10955 .await
10956 .expect("session lock")
10957 .to_messages_for_current_path();
10958 assert!(
10959 stored.iter().any(|message| matches!(
10960 message,
10961 Message::Custom(CustomMessage { custom_type, details, display: false, .. })
10962 if custom_type == SEMANTIC_CONTEXT_CUSTOM_TYPE
10963 && details
10964 .as_ref()
10965 .and_then(|value| value.get("bundleRevision"))
10966 .and_then(Value::as_str)
10967 == Some(revision.as_str())
10968 )),
10969 "hidden semantic context provenance was not persisted in session messages: {stored:?}"
10970 );
10971 assert_eq!(agent_session.agent.system_prompt(), Some("base prompt"));
10972 });
10973 }
10974
10975 #[test]
10976 fn enable_extensions_policy_resolution_defaults_to_permissive() {
10977 let policy = AgentSession::resolve_extension_policy_for_enable(None, None);
10978 assert_eq!(
10979 policy.mode,
10980 crate::extensions::ExtensionPolicyMode::Permissive
10981 );
10982 }
10983
10984 #[test]
10985 fn enable_extensions_policy_resolution_respects_config_default_toggle() {
10986 let config = crate::config::Config {
10987 extension_policy: Some(crate::config::ExtensionPolicyConfig {
10988 profile: None,
10989 default_permissive: Some(false),
10990 allow_dangerous: None,
10991 }),
10992 ..Default::default()
10993 };
10994 let policy = AgentSession::resolve_extension_policy_for_enable(Some(&config), None);
10995 assert_eq!(policy.mode, crate::extensions::ExtensionPolicyMode::Strict);
10996 }
10997
10998 #[test]
10999 fn enable_extensions_policy_resolution_prefers_explicit_policy() {
11000 let config = crate::config::Config {
11001 extension_policy: Some(crate::config::ExtensionPolicyConfig {
11002 profile: None,
11003 default_permissive: Some(false),
11004 allow_dangerous: None,
11005 }),
11006 ..Default::default()
11007 };
11008 let explicit = crate::extensions::PolicyProfile::Permissive.to_policy();
11009 let policy =
11010 AgentSession::resolve_extension_policy_for_enable(Some(&config), Some(explicit));
11011 assert_eq!(
11012 policy.mode,
11013 crate::extensions::ExtensionPolicyMode::Permissive
11014 );
11015 }
11016
11017 #[test]
11018 fn test_extract_tool_calls() {
11019 let content = vec![
11020 ContentBlock::Text(TextContent::new("Hello")),
11021 ContentBlock::ToolCall(ToolCall {
11022 id: "tc1".to_string(),
11023 name: "read".to_string(),
11024 arguments: serde_json::json!({"path": "file.txt"}),
11025 thought_signature: None,
11026 }),
11027 ContentBlock::Text(TextContent::new("World")),
11028 ContentBlock::ToolCall(ToolCall {
11029 id: "tc2".to_string(),
11030 name: "bash".to_string(),
11031 arguments: serde_json::json!({"command": "ls"}),
11032 thought_signature: None,
11033 }),
11034 ];
11035
11036 let tool_calls = extract_tool_calls(&content);
11037 assert_eq!(tool_calls.len(), 2);
11038 assert_eq!(tool_calls[0].name, "read");
11039 assert_eq!(tool_calls[1].name, "bash");
11040 }
11041
11042 #[test]
11043 fn test_agent_config_default() {
11044 let config = AgentConfig::default();
11051 let expected = resolved_max_tool_iterations_default();
11052 assert_eq!(config.max_tool_iterations, expected);
11053 assert!(config.system_prompt.is_none());
11054 assert!(!config.block_images);
11055 }
11056
11057 #[test]
11058 fn resolve_max_tool_iterations_handles_unset_empty_and_whitespace() {
11059 assert_eq!(
11060 resolve_max_tool_iterations(None),
11061 MAX_TOOL_ITERATIONS_DEFAULT
11062 );
11063 assert_eq!(
11064 resolve_max_tool_iterations(Some("")),
11065 MAX_TOOL_ITERATIONS_DEFAULT
11066 );
11067 assert_eq!(
11068 resolve_max_tool_iterations(Some(" ")),
11069 MAX_TOOL_ITERATIONS_DEFAULT
11070 );
11071 }
11072
11073 #[test]
11074 fn resolve_max_tool_iterations_rejects_zero_and_invalid() {
11075 assert_eq!(
11076 resolve_max_tool_iterations(Some("0")),
11077 MAX_TOOL_ITERATIONS_DEFAULT
11078 );
11079 assert_eq!(
11080 resolve_max_tool_iterations(Some("not-a-number")),
11081 MAX_TOOL_ITERATIONS_DEFAULT
11082 );
11083 assert_eq!(
11084 resolve_max_tool_iterations(Some("-5")),
11085 MAX_TOOL_ITERATIONS_DEFAULT
11086 );
11087 assert_eq!(
11088 resolve_max_tool_iterations(Some("3.14")),
11089 MAX_TOOL_ITERATIONS_DEFAULT
11090 );
11091 }
11092
11093 #[test]
11094 fn resolve_max_tool_iterations_accepts_valid_overrides_and_trims_whitespace() {
11095 assert_eq!(resolve_max_tool_iterations(Some("1")), 1);
11096 assert_eq!(resolve_max_tool_iterations(Some("100")), 100);
11097 assert_eq!(resolve_max_tool_iterations(Some(" 200 ")), 200);
11098 assert_eq!(resolve_max_tool_iterations(Some("999")), 999);
11099 }
11100
11101 #[test]
11102 fn resolve_max_tool_iterations_clamps_above_ceiling() {
11103 assert_eq!(
11104 resolve_max_tool_iterations(Some("99999")),
11105 MAX_TOOL_ITERATIONS_CEILING
11106 );
11107 assert_eq!(
11109 resolve_max_tool_iterations(Some("1000")),
11110 MAX_TOOL_ITERATIONS_CEILING
11111 );
11112 }
11113
11114 #[test]
11115 fn clamp_max_tool_iterations_matches_resolve_semantics() {
11116 assert_eq!(clamp_max_tool_iterations(None), MAX_TOOL_ITERATIONS_DEFAULT);
11118 assert_eq!(
11119 clamp_max_tool_iterations(Some(0)),
11120 MAX_TOOL_ITERATIONS_DEFAULT
11121 );
11122 assert_eq!(clamp_max_tool_iterations(Some(7)), 7);
11123 assert_eq!(
11124 clamp_max_tool_iterations(Some(usize::MAX)),
11125 MAX_TOOL_ITERATIONS_CEILING
11126 );
11127 }
11128
11129 #[test]
11130 fn iteration_warning_fires_at_80_percent_for_default_cap() {
11131 assert!(!should_warn_at_iteration_threshold(39, 50));
11133 assert!(should_warn_at_iteration_threshold(40, 50));
11134 assert!(should_warn_at_iteration_threshold(50, 50));
11135 assert!(!should_warn_at_iteration_threshold(0, 50));
11137 }
11138
11139 #[test]
11140 fn iteration_warning_fires_at_80_percent_for_custom_caps() {
11141 for (cap, threshold) in [(100usize, 80usize), (200, 160), (1000, 800)] {
11142 assert!(
11143 !should_warn_at_iteration_threshold(threshold - 1, cap),
11144 "expected no warning below threshold (current=cap={cap}, threshold={threshold})"
11145 );
11146 assert!(
11147 should_warn_at_iteration_threshold(threshold, cap),
11148 "expected warning at threshold (cap={cap}, threshold={threshold})"
11149 );
11150 }
11151 }
11152
11153 #[test]
11154 fn iteration_warning_skipped_for_caps_below_minimum() {
11155 for cap in 0..ITERATION_WARN_MIN_CAP {
11159 for current in 0..=cap.saturating_add(2) {
11160 assert!(
11161 !should_warn_at_iteration_threshold(current, cap),
11162 "should not warn at current={current} cap={cap}"
11163 );
11164 }
11165 }
11166 }
11167
11168 #[test]
11169 fn iteration_warning_handles_minimum_warnable_cap_boundary() {
11170 assert!(!should_warn_at_iteration_threshold(3, 5));
11172 assert!(should_warn_at_iteration_threshold(4, 5));
11173 assert!(should_warn_at_iteration_threshold(5, 5));
11174 }
11175
11176 #[test]
11177 fn iteration_warning_handles_overflow_resistant_caps() {
11178 assert!(!should_warn_at_iteration_threshold(1_000_000, usize::MAX));
11185 assert!(!should_warn_at_iteration_threshold(
11186 usize::MAX / 6,
11187 usize::MAX
11188 ));
11189 assert!(should_warn_at_iteration_threshold(
11191 usize::MAX / 5,
11192 usize::MAX
11193 ));
11194 }
11195
11196 #[test]
11197 fn iteration_handoff_steering_text_is_self_describing() {
11198 let text = iteration_handoff_steering_text(42, 50);
11203 assert!(text.contains("[runtime]"));
11204 assert!(text.contains("Tool-iteration budget at >=80%"));
11205 assert!(text.contains("used 42 of 50"));
11206 assert!(text.contains("graceful handoff"));
11207 assert!(text.contains("incomplete-handoff"));
11208 assert!(text.contains("Do NOT compress"));
11209 }
11210
11211 #[test]
11212 fn filter_image_blocks_replaces_images_with_deduped_placeholder_text() {
11213 let mut blocks = vec![
11214 sample_image_block(),
11215 sample_image_block(),
11216 ContentBlock::Text(TextContent::new("tail")),
11217 sample_image_block(),
11218 ];
11219
11220 let removed = filter_image_blocks(&mut blocks);
11221
11222 assert_eq!(removed, 3);
11223 assert!(
11224 !blocks
11225 .iter()
11226 .any(|block| matches!(block, ContentBlock::Image(_)))
11227 );
11228 assert!(matches!(
11229 blocks.first(),
11230 Some(ContentBlock::Text(TextContent { text, .. }))
11231 if text.as_str().eq(BLOCK_IMAGES_PLACEHOLDER)
11232 ));
11233 assert!(matches!(
11234 blocks.get(1),
11235 Some(ContentBlock::Text(TextContent { text, .. })) if text.as_str().eq("tail")
11236 ));
11237 assert!(matches!(
11238 blocks.get(2),
11239 Some(ContentBlock::Text(TextContent { text, .. }))
11240 if text.as_str().eq(BLOCK_IMAGES_PLACEHOLDER)
11241 ));
11242 }
11243
11244 #[test]
11245 fn filter_images_for_provider_filters_images_from_all_block_message_types() {
11246 let mut messages = vec![
11247 Message::User(UserMessage {
11248 content: UserContent::Blocks(vec![
11249 ContentBlock::Text(TextContent::new("hello")),
11250 sample_image_block(),
11251 ]),
11252 timestamp: 0,
11253 }),
11254 Message::Assistant(Arc::new(AssistantMessage {
11255 content: vec![sample_image_block()],
11256 api: "test".to_string(),
11257 provider: "test".to_string(),
11258 model: "test".to_string(),
11259 usage: Usage::default(),
11260 stop_reason: StopReason::Stop,
11261 error_message: None,
11262 timestamp: 0,
11263 })),
11264 Message::tool_result(ToolResultMessage {
11265 tool_call_id: "tc1".to_string(),
11266 tool_name: "read".to_string(),
11267 content: vec![
11268 sample_image_block(),
11269 ContentBlock::Text(TextContent::new("ok")),
11270 ],
11271 details: None,
11272 is_error: false,
11273 timestamp: 0,
11274 }),
11275 ];
11276
11277 let stats = filter_images_for_provider(&mut messages);
11278
11279 assert_eq!(stats.removed_images, 3);
11280 assert_eq!(stats.affected_messages, 3);
11281 assert_eq!(
11282 messages.iter().map(image_count_in_message).sum::<usize>(),
11283 0,
11284 "no images should remain in provider-bound context"
11285 );
11286 }
11287
11288 #[test]
11289 fn build_context_strips_images_when_block_images_enabled() {
11290 let mut agent = Agent::new(
11291 Arc::new(SilentProvider),
11292 ToolRegistry::new(&[], Path::new("."), None),
11293 AgentConfig {
11294 system_prompt: None,
11295 max_tool_iterations: 50,
11296 stream_options: StreamOptions::default(),
11297 block_images: true,
11298 fail_closed_hooks: false,
11299 tool_approval: None,
11300 },
11301 );
11302 agent.add_message(Message::User(UserMessage {
11303 content: UserContent::Blocks(vec![sample_image_block()]),
11304 timestamp: 0,
11305 }));
11306
11307 let context = agent.build_context();
11308 assert_eq!(context.messages.len(), 1);
11309 assert_eq!(image_count_in_message(&context.messages[0]), 0);
11310 assert!(matches!(
11311 &context.messages[0],
11312 Message::User(UserMessage {
11313 content: UserContent::Blocks(blocks),
11314 ..
11315 }) if blocks
11316 .iter()
11317 .any(|block| matches!(block, ContentBlock::Text(TextContent { text, .. }) if text.as_str().eq(BLOCK_IMAGES_PLACEHOLDER)))
11318 ));
11319 }
11320
11321 #[test]
11322 fn build_context_keeps_images_when_block_images_disabled() {
11323 let mut agent = Agent::new(
11324 Arc::new(SilentProvider),
11325 ToolRegistry::new(&[], Path::new("."), None),
11326 AgentConfig {
11327 system_prompt: None,
11328 max_tool_iterations: 50,
11329 stream_options: StreamOptions::default(),
11330 block_images: false,
11331 fail_closed_hooks: false,
11332 tool_approval: None,
11333 },
11334 );
11335 agent.add_message(Message::User(UserMessage {
11336 content: UserContent::Blocks(vec![sample_image_block()]),
11337 timestamp: 0,
11338 }));
11339
11340 let context = agent.build_context();
11341 assert_eq!(context.messages.len(), 1);
11342 assert_eq!(image_count_in_message(&context.messages[0]), 1);
11343 }
11344
11345 #[test]
11346 fn auto_compaction_start_serializes_with_pi_mono_compatible_type_tag() {
11347 let event = AgentEvent::AutoCompactionStart {
11348 reason: "threshold".to_string(),
11349 };
11350 let json = serde_json::to_value(&event).unwrap();
11351 assert_eq!(json["type"], "auto_compaction_start");
11352 assert_eq!(json["reason"], "threshold");
11353 }
11354
11355 #[test]
11356 fn auto_compaction_end_serializes_with_pi_mono_compatible_fields() {
11357 let event = AgentEvent::AutoCompactionEnd {
11358 result: Some(serde_json::json!({"tokens_before": 5000, "tokens_after": 2000})),
11359 aborted: false,
11360 will_retry: false,
11361 error_message: None,
11362 };
11363 let json = serde_json::to_value(&event).unwrap();
11364 assert_eq!(json["type"], "auto_compaction_end");
11365 assert_eq!(json["aborted"], false);
11366 assert_eq!(json["willRetry"], false);
11367 assert!(json.get("errorMessage").is_none()); assert!(json["result"].is_object());
11369 }
11370
11371 #[test]
11372 fn auto_compaction_end_includes_error_message_when_present() {
11373 let event = AgentEvent::AutoCompactionEnd {
11374 result: None,
11375 aborted: true,
11376 will_retry: false,
11377 error_message: Some("Compaction failed".to_string()),
11378 };
11379 let json = serde_json::to_value(&event).unwrap();
11380 assert_eq!(json["type"], "auto_compaction_end");
11381 assert_eq!(json["aborted"], true);
11382 assert_eq!(json["errorMessage"], "Compaction failed");
11383 }
11384
11385 #[test]
11386 fn auto_retry_start_serializes_with_camel_case_fields() {
11387 let event = AgentEvent::AutoRetryStart {
11388 attempt: 1,
11389 max_attempts: 3,
11390 delay_ms: 2000,
11391 error_message: "Rate limited".to_string(),
11392 };
11393 let json = serde_json::to_value(&event).unwrap();
11394 assert_eq!(json["type"], "auto_retry_start");
11395 assert_eq!(json["attempt"], 1);
11396 assert_eq!(json["maxAttempts"], 3);
11397 assert_eq!(json["delayMs"], 2000);
11398 assert_eq!(json["errorMessage"], "Rate limited");
11399 }
11400
11401 #[test]
11402 fn auto_retry_end_serializes_success_and_omits_null_final_error() {
11403 let event = AgentEvent::AutoRetryEnd {
11404 success: true,
11405 attempt: 2,
11406 final_error: None,
11407 };
11408 let json = serde_json::to_value(&event).unwrap();
11409 assert_eq!(json["type"], "auto_retry_end");
11410 assert_eq!(json["success"], true);
11411 assert_eq!(json["attempt"], 2);
11412 assert!(json.get("finalError").is_none());
11413 }
11414
11415 #[test]
11416 fn auto_retry_end_includes_final_error_on_failure() {
11417 let event = AgentEvent::AutoRetryEnd {
11418 success: false,
11419 attempt: 3,
11420 final_error: Some("Max retries exceeded".to_string()),
11421 };
11422 let json = serde_json::to_value(&event).unwrap();
11423 assert_eq!(json["type"], "auto_retry_end");
11424 assert_eq!(json["success"], false);
11425 assert_eq!(json["attempt"], 3);
11426 assert_eq!(json["finalError"], "Max retries exceeded");
11427 }
11428
11429 #[test]
11430 fn message_queue_push_increments_seq_and_counts_both_queues() {
11431 let mut queue = MessageQueue::new(QueueMode::OneAtATime, QueueMode::OneAtATime);
11432 assert_eq!(queue.pending_count(), 0);
11433
11434 assert_eq!(queue.push_steering(user_message("s1")), 0);
11435 assert_eq!(queue.push_follow_up(user_message("f1")), 1);
11436 assert_eq!(queue.push_steering(user_message("s2")), 2);
11437
11438 assert_eq!(queue.pending_count(), 3);
11439 }
11440
11441 #[test]
11442 fn message_queue_pop_steering_one_at_a_time_preserves_order() {
11443 let mut queue = MessageQueue::new(QueueMode::OneAtATime, QueueMode::OneAtATime);
11444 queue.push_steering(user_message("s1"));
11445 queue.push_steering(user_message("s2"));
11446
11447 let first = queue.pop_steering();
11448 assert_eq!(first.len(), 1);
11449 assert_user_text(&first[0], "s1");
11450 assert_eq!(queue.pending_count(), 1);
11451
11452 let second = queue.pop_steering();
11453 assert_eq!(second.len(), 1);
11454 assert_user_text(&second[0], "s2");
11455 assert_eq!(queue.pending_count(), 0);
11456
11457 let empty = queue.pop_steering();
11458 assert!(empty.is_empty());
11459 }
11460
11461 #[test]
11462 fn message_queue_pop_respects_queue_modes_per_kind() {
11463 let mut queue = MessageQueue::new(QueueMode::All, QueueMode::OneAtATime);
11464 queue.push_steering(user_message("s1"));
11465 queue.push_steering(user_message("s2"));
11466 queue.push_follow_up(user_message("f1"));
11467 queue.push_follow_up(user_message("f2"));
11468
11469 let steering = queue.pop_steering();
11470 assert_eq!(steering.len(), 2);
11471 assert_user_text(&steering[0], "s1");
11472 assert_user_text(&steering[1], "s2");
11473 assert_eq!(queue.pending_count(), 2);
11474
11475 let follow_up = queue.pop_follow_up();
11476 assert_eq!(follow_up.len(), 1);
11477 assert_user_text(&follow_up[0], "f1");
11478 assert_eq!(queue.pending_count(), 1);
11479
11480 let follow_up = queue.pop_follow_up();
11481 assert_eq!(follow_up.len(), 1);
11482 assert_user_text(&follow_up[0], "f2");
11483 assert_eq!(queue.pending_count(), 0);
11484 }
11485
11486 #[test]
11487 fn message_queue_set_modes_applies_to_existing_messages() {
11488 let mut queue = MessageQueue::new(QueueMode::OneAtATime, QueueMode::OneAtATime);
11489 queue.push_steering(user_message("s1"));
11490 queue.push_steering(user_message("s2"));
11491
11492 let first = queue.pop_steering();
11493 assert_eq!(first.len(), 1);
11494 assert_user_text(&first[0], "s1");
11495
11496 queue.set_modes(QueueMode::All, QueueMode::OneAtATime);
11497 let remaining = queue.pop_steering();
11498 assert_eq!(remaining.len(), 1);
11499 assert_user_text(&remaining[0], "s2");
11500 }
11501
11502 fn build_switch_test_session(auth: &AuthStorage) -> AgentSession {
11503 let registry = ModelRegistry::load(auth, None);
11504 let current_entry = registry
11505 .find("anthropic", "claude-sonnet-4-5")
11506 .expect("anthropic model in registry");
11507 let provider = crate::providers::create_provider(¤t_entry, None)
11508 .expect("create anthropic provider");
11509 let tools = ToolRegistry::new(&[], Path::new("."), None);
11510 let mut stream_options = StreamOptions {
11511 api_key: Some("stale-key".to_string()),
11512 ..Default::default()
11513 };
11514 let _ = stream_options
11515 .headers
11516 .insert("x-stale-header".to_string(), "stale-value".to_string());
11517 let agent = Agent::new(
11518 provider,
11519 tools,
11520 AgentConfig {
11521 system_prompt: None,
11522 max_tool_iterations: 50,
11523 stream_options,
11524 block_images: false,
11525 fail_closed_hooks: false,
11526 tool_approval: None,
11527 },
11528 );
11529
11530 let mut session = Session::in_memory();
11531 session.header.provider = Some("anthropic".to_string());
11532 session.header.model_id = Some("claude-sonnet-4-5".to_string());
11533
11534 let mut agent_session = AgentSession::new(
11535 agent,
11536 Arc::new(Mutex::new(session)),
11537 false,
11538 ResolvedCompactionSettings::default(),
11539 );
11540 agent_session.set_model_registry(registry);
11541 agent_session.set_auth_storage(auth.clone());
11542 agent_session
11543 }
11544
11545 #[test]
11546 fn compaction_runtime_handle_creates_fallback_runtime() {
11547 let dir = tempfile::tempdir().expect("tempdir");
11548 let auth_path = dir.path().join("auth.json");
11549 let auth = AuthStorage::load(auth_path).expect("load auth");
11550 let mut agent_session = build_switch_test_session(&auth);
11551
11552 assert!(agent_session.compaction_runtime.is_none());
11553 assert!(agent_session.runtime_handle.is_none());
11554
11555 let runtime_handle = agent_session
11556 .compaction_runtime_handle()
11557 .expect("create fallback compaction runtime");
11558 let join = runtime_handle.spawn(async { 7_u8 });
11559 assert_eq!(futures::executor::block_on(join), 7);
11560
11561 assert!(agent_session.compaction_runtime.is_some());
11562 assert!(agent_session.runtime_handle.is_some());
11563 }
11564
11565 #[test]
11566 fn apply_session_model_selection_updates_stream_credentials_and_headers() {
11567 let dir = tempfile::tempdir().expect("tempdir");
11568 let auth_path = dir.path().join("auth.json");
11569 let mut auth = AuthStorage::load(auth_path).expect("load auth");
11570 auth.set(
11571 "anthropic",
11572 AuthCredential::ApiKey {
11573 key: "anthropic-key".to_string(),
11574 },
11575 );
11576 auth.set(
11577 "openai",
11578 AuthCredential::ApiKey {
11579 key: "openai-key".to_string(),
11580 },
11581 );
11582
11583 let mut agent_session = build_switch_test_session(&auth);
11584 agent_session
11585 .apply_session_model_selection("openai", "gpt-4o")
11586 .expect("switch should update stream options");
11587
11588 assert_eq!(agent_session.agent.provider().name(), "openai");
11589 assert_eq!(agent_session.agent.provider().model_id(), "gpt-4o");
11590 assert_eq!(
11591 agent_session.agent.stream_options().api_key.as_deref(),
11592 Some("openai-key")
11593 );
11594 assert!(
11595 agent_session.agent.stream_options().headers.is_empty(),
11596 "stream headers should be refreshed from selected model entry"
11597 );
11598 }
11599
11600 #[test]
11601 fn apply_session_model_selection_clears_stale_key_for_keyless_target() {
11602 let dir = tempfile::tempdir().expect("tempdir");
11603 let auth_path = dir.path().join("auth.json");
11604 let mut auth = AuthStorage::load(auth_path).expect("load auth");
11605 auth.set(
11606 "anthropic",
11607 AuthCredential::ApiKey {
11608 key: "anthropic-key".to_string(),
11609 },
11610 );
11611
11612 let mut registry = ModelRegistry::load(&auth, None);
11613 registry.merge_entries(vec![ModelEntry {
11614 model: Model {
11615 id: "local-model".to_string(),
11616 name: "Local Model".to_string(),
11617 api: "openai-completions".to_string(),
11618 provider: "acme-local".to_string(),
11619 base_url: "https://example.invalid/v1".to_string(),
11620 reasoning: true,
11621 input: vec![InputType::Text],
11622 cost: ModelCost {
11623 input: 0.0,
11624 output: 0.0,
11625 cache_read: 0.0,
11626 cache_write: 0.0,
11627 },
11628 context_window: 128_000,
11629 max_tokens: 8_192,
11630 headers: HashMap::new(),
11631 },
11632 api_key: None,
11633 headers: HashMap::new(),
11634 auth_header: false,
11635 compat: None,
11636 oauth_config: None,
11637 }]);
11638
11639 let mut agent_session = build_switch_test_session(&auth);
11640 agent_session.set_model_registry(registry);
11641 agent_session
11642 .apply_session_model_selection("acme-local", "local-model")
11643 .expect("keyless local model should still activate");
11644
11645 assert_eq!(agent_session.agent.provider().name(), "acme-local");
11646 assert_eq!(
11647 agent_session.agent.stream_options().api_key,
11648 None,
11649 "stale key must be cleared when target model has no configured key"
11650 );
11651 }
11652
11653 #[test]
11654 fn apply_session_model_selection_treats_blank_model_key_as_missing_credential() {
11655 let dir = tempfile::tempdir().expect("tempdir");
11656 let auth_path = dir.path().join("auth.json");
11657 let auth = AuthStorage::load(auth_path).expect("load auth");
11658
11659 let mut registry = ModelRegistry::load(&auth, None);
11660 registry.merge_entries(vec![ModelEntry {
11661 model: Model {
11662 id: "blank-model".to_string(),
11663 name: "Blank Model".to_string(),
11664 api: "openai-completions".to_string(),
11665 provider: "acme".to_string(),
11666 base_url: "https://example.invalid/v1".to_string(),
11667 reasoning: true,
11668 input: vec![InputType::Text],
11669 cost: ModelCost {
11670 input: 0.0,
11671 output: 0.0,
11672 cache_read: 0.0,
11673 cache_write: 0.0,
11674 },
11675 context_window: 128_000,
11676 max_tokens: 8_192,
11677 headers: HashMap::new(),
11678 },
11679 api_key: Some(" ".to_string()),
11680 headers: HashMap::new(),
11681 auth_header: true,
11682 compat: None,
11683 oauth_config: None,
11684 }]);
11685
11686 let mut agent_session = build_switch_test_session(&auth);
11687 agent_session.set_model_registry(registry);
11688 let err = agent_session
11689 .apply_session_model_selection("acme", "blank-model")
11690 .expect_err("blank keys must not satisfy credential requirements");
11691
11692 assert!(
11693 err.to_string()
11694 .contains("Missing credentials for acme/blank-model"),
11695 "unexpected error: {err}"
11696 );
11697 assert_eq!(agent_session.agent.provider().name(), "anthropic");
11698 assert_eq!(
11699 agent_session.agent.stream_options().api_key,
11700 Some("stale-key".to_string()),
11701 "failed switches must preserve the prior runtime credentials"
11702 );
11703 }
11704
11705 #[test]
11706 fn set_provider_model_preserves_session_header_when_switch_fails() {
11707 let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
11708 .build()
11709 .expect("build runtime");
11710
11711 runtime.block_on(async {
11712 let dir = tempfile::tempdir().expect("tempdir");
11713 let auth_path = dir.path().join("auth.json");
11714 let auth = AuthStorage::load(auth_path).expect("load auth");
11715 let mut agent_session = build_switch_test_session(&auth);
11716
11717 {
11718 let cx = crate::agent_cx::AgentCx::for_request();
11719 let mut session = agent_session
11720 .session
11721 .lock(cx.cx())
11722 .await
11723 .expect("session lock");
11724 session.header.provider = Some("anthropic".to_string());
11725 session.header.model_id = Some("claude-sonnet-4-5".to_string());
11726 }
11727
11728 let err = agent_session
11729 .set_provider_model("missing-provider", "missing-model")
11730 .await
11731 .expect_err("missing model should not switch");
11732 assert!(
11733 err.to_string()
11734 .contains("Unable to switch provider/model to missing-provider/missing-model"),
11735 "unexpected error: {err}"
11736 );
11737 assert_eq!(agent_session.agent.provider().name(), "anthropic");
11738 assert_eq!(
11739 agent_session.agent.provider().model_id(),
11740 "claude-sonnet-4-5"
11741 );
11742
11743 let cx = crate::agent_cx::AgentCx::for_request();
11744 let session = agent_session
11745 .session
11746 .lock(cx.cx())
11747 .await
11748 .expect("session lock");
11749 assert_eq!(session.header.provider.as_deref(), Some("anthropic"));
11750 assert_eq!(
11751 session.header.model_id.as_deref(),
11752 Some("claude-sonnet-4-5")
11753 );
11754 });
11755 }
11756
11757 #[test]
11758 fn set_provider_model_rejects_missing_credentials_without_switching() {
11759 let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
11760 .build()
11761 .expect("build runtime");
11762
11763 runtime.block_on(async {
11764 let dir = tempfile::tempdir().expect("tempdir");
11765 let auth_path = dir.path().join("auth.json");
11766 let auth = AuthStorage::load(auth_path).expect("load auth");
11767 let mut agent_session = build_switch_test_session(&auth);
11768
11769 {
11770 let cx = crate::agent_cx::AgentCx::for_request();
11771 let mut session = agent_session
11772 .session
11773 .lock(cx.cx())
11774 .await
11775 .expect("session lock");
11776 session.header.provider = Some("anthropic".to_string());
11777 session.header.model_id = Some("claude-sonnet-4-5".to_string());
11778 }
11779
11780 let err = agent_session
11781 .set_provider_model("openai", "gpt-4o")
11782 .await
11783 .expect_err("missing credentials should abort model switch");
11784 assert!(
11785 err.to_string()
11786 .contains("Missing credentials for openai/gpt-4o"),
11787 "unexpected error: {err}"
11788 );
11789 assert_eq!(agent_session.agent.provider().name(), "anthropic");
11790 assert_eq!(
11791 agent_session.agent.provider().model_id(),
11792 "claude-sonnet-4-5"
11793 );
11794
11795 let cx = crate::agent_cx::AgentCx::for_request();
11796 let session = agent_session
11797 .session
11798 .lock(cx.cx())
11799 .await
11800 .expect("session lock");
11801 assert_eq!(session.header.provider.as_deref(), Some("anthropic"));
11802 assert_eq!(
11803 session.header.model_id.as_deref(),
11804 Some("claude-sonnet-4-5")
11805 );
11806 });
11807 }
11808
11809 #[test]
11810 fn set_provider_model_clamps_thinking_for_non_reasoning_targets() {
11811 let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
11812 .build()
11813 .expect("build runtime");
11814
11815 runtime.block_on(async {
11816 let dir = tempfile::tempdir().expect("tempdir");
11817 let auth_path = dir.path().join("auth.json");
11818 let auth = AuthStorage::load(auth_path).expect("load auth");
11819
11820 let mut registry = ModelRegistry::load(&auth, None);
11821 registry.merge_entries(vec![ModelEntry {
11822 model: Model {
11823 id: "plain-model".to_string(),
11824 name: "Plain Model".to_string(),
11825 api: "openai-completions".to_string(),
11826 provider: "acme".to_string(),
11827 base_url: "https://example.invalid/v1".to_string(),
11828 reasoning: false,
11829 input: vec![InputType::Text],
11830 cost: ModelCost {
11831 input: 0.0,
11832 output: 0.0,
11833 cache_read: 0.0,
11834 cache_write: 0.0,
11835 },
11836 context_window: 128_000,
11837 max_tokens: 8_192,
11838 headers: HashMap::new(),
11839 },
11840 api_key: None,
11841 headers: HashMap::new(),
11842 auth_header: false,
11843 compat: None,
11844 oauth_config: None,
11845 }]);
11846
11847 let mut agent_session = build_switch_test_session(&auth);
11848 agent_session.set_model_registry(registry);
11849 agent_session.agent.stream_options_mut().thinking_level =
11850 Some(crate::model::ThinkingLevel::High);
11851
11852 {
11853 let cx = crate::agent_cx::AgentCx::for_request();
11854 let mut session = agent_session
11855 .session
11856 .lock(cx.cx())
11857 .await
11858 .expect("session lock");
11859 session.header.thinking_level = Some("high".to_string());
11860 }
11861
11862 agent_session
11863 .set_provider_model("acme", "plain-model")
11864 .await
11865 .expect("switch should clamp unsupported thinking");
11866
11867 assert_eq!(agent_session.agent.provider().name(), "acme");
11868 assert_eq!(agent_session.agent.provider().model_id(), "plain-model");
11869 assert_eq!(
11870 agent_session.agent.stream_options().thinking_level,
11871 Some(crate::model::ThinkingLevel::Off)
11872 );
11873
11874 let cx = crate::agent_cx::AgentCx::for_request();
11875 let session = agent_session
11876 .session
11877 .lock(cx.cx())
11878 .await
11879 .expect("session lock");
11880 assert_eq!(session.header.provider.as_deref(), Some("acme"));
11881 assert_eq!(session.header.model_id.as_deref(), Some("plain-model"));
11882 assert_eq!(session.header.thinking_level.as_deref(), Some("off"));
11883 });
11884 }
11885
11886 #[test]
11887 fn set_provider_model_records_model_change_once() {
11888 let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
11889 .build()
11890 .expect("build runtime");
11891
11892 runtime.block_on(async {
11893 let dir = tempfile::tempdir().expect("tempdir");
11894 let auth_path = dir.path().join("auth.json");
11895 let mut auth = AuthStorage::load(auth_path).expect("load auth");
11896 auth.set(
11897 "anthropic",
11898 AuthCredential::ApiKey {
11899 key: "anthropic-key".to_string(),
11900 },
11901 );
11902 auth.set(
11903 "openai",
11904 AuthCredential::ApiKey {
11905 key: "openai-key".to_string(),
11906 },
11907 );
11908
11909 let mut agent_session = build_switch_test_session(&auth);
11910 agent_session
11911 .set_provider_model("openai", "gpt-4o")
11912 .await
11913 .expect("switch model");
11914 agent_session
11915 .set_provider_model("openai", "gpt-4o")
11916 .await
11917 .expect("repeat same model");
11918
11919 let cx = crate::agent_cx::AgentCx::for_request();
11920 let session = agent_session
11921 .session
11922 .lock(cx.cx())
11923 .await
11924 .expect("session lock");
11925 let model_changes = session
11926 .entries_for_current_path()
11927 .iter()
11928 .filter(|entry| matches!(entry, crate::session::SessionEntry::ModelChange(_)))
11929 .count();
11930 assert_eq!(model_changes, 1);
11931 });
11932 }
11933
11934 #[test]
11935 fn sync_runtime_selection_from_session_header_clamps_and_normalizes_thinking() {
11936 let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
11937 .build()
11938 .expect("build runtime");
11939
11940 runtime.block_on(async {
11941 let dir = tempfile::tempdir().expect("tempdir");
11942 let auth_path = dir.path().join("auth.json");
11943 let auth = AuthStorage::load(auth_path).expect("load auth");
11944
11945 let mut registry = ModelRegistry::load(&auth, None);
11946 registry.merge_entries(vec![ModelEntry {
11947 model: Model {
11948 id: "plain-model".to_string(),
11949 name: "Plain Model".to_string(),
11950 api: "openai-completions".to_string(),
11951 provider: "acme".to_string(),
11952 base_url: "https://example.invalid/v1".to_string(),
11953 reasoning: false,
11954 input: vec![InputType::Text],
11955 cost: ModelCost {
11956 input: 0.0,
11957 output: 0.0,
11958 cache_read: 0.0,
11959 cache_write: 0.0,
11960 },
11961 context_window: 128_000,
11962 max_tokens: 8_192,
11963 headers: HashMap::new(),
11964 },
11965 api_key: None,
11966 headers: HashMap::new(),
11967 auth_header: false,
11968 compat: None,
11969 oauth_config: None,
11970 }]);
11971
11972 let mut agent_session = build_switch_test_session(&auth);
11973 agent_session.set_model_registry(registry);
11974 agent_session.agent.stream_options_mut().thinking_level =
11975 Some(crate::model::ThinkingLevel::High);
11976
11977 {
11978 let cx = crate::agent_cx::AgentCx::for_request();
11979 let mut session = agent_session
11980 .session
11981 .lock(cx.cx())
11982 .await
11983 .expect("session lock");
11984 session.header.provider = Some("acme".to_string());
11985 session.header.model_id = Some("plain-model".to_string());
11986 session.header.thinking_level = Some("high".to_string());
11987 }
11988
11989 agent_session
11990 .sync_runtime_selection_from_session_header()
11991 .await
11992 .expect("sync runtime selection");
11993
11994 assert_eq!(agent_session.agent.provider().name(), "acme");
11995 assert_eq!(agent_session.agent.provider().model_id(), "plain-model");
11996 assert_eq!(
11997 agent_session.agent.stream_options().thinking_level,
11998 Some(crate::model::ThinkingLevel::Off)
11999 );
12000
12001 let cx = crate::agent_cx::AgentCx::for_request();
12002 let session = agent_session
12003 .session
12004 .lock(cx.cx())
12005 .await
12006 .expect("session lock");
12007 assert_eq!(session.header.thinking_level.as_deref(), Some("off"));
12008 let thinking_changes = session
12009 .entries_for_current_path()
12010 .iter()
12011 .filter(|entry| {
12012 matches!(entry, crate::session::SessionEntry::ThinkingLevelChange(_))
12013 })
12014 .count();
12015 assert_eq!(thinking_changes, 1);
12016 });
12017 }
12018
12019 #[test]
12020 fn sync_runtime_selection_from_session_header_clamps_current_thinking_when_header_omits_it() {
12021 let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
12022 .build()
12023 .expect("build runtime");
12024
12025 runtime.block_on(async {
12026 let dir = tempfile::tempdir().expect("tempdir");
12027 let auth_path = dir.path().join("auth.json");
12028 let auth = AuthStorage::load(auth_path).expect("load auth");
12029
12030 let mut registry = ModelRegistry::load(&auth, None);
12031 registry.merge_entries(vec![ModelEntry {
12032 model: Model {
12033 id: "plain-model".to_string(),
12034 name: "Plain Model".to_string(),
12035 api: "openai-completions".to_string(),
12036 provider: "acme".to_string(),
12037 base_url: "https://example.invalid/v1".to_string(),
12038 reasoning: false,
12039 input: vec![InputType::Text],
12040 cost: ModelCost {
12041 input: 0.0,
12042 output: 0.0,
12043 cache_read: 0.0,
12044 cache_write: 0.0,
12045 },
12046 context_window: 128_000,
12047 max_tokens: 8_192,
12048 headers: HashMap::new(),
12049 },
12050 api_key: None,
12051 headers: HashMap::new(),
12052 auth_header: false,
12053 compat: None,
12054 oauth_config: None,
12055 }]);
12056
12057 let mut agent_session = build_switch_test_session(&auth);
12058 agent_session.set_model_registry(registry);
12059 agent_session.agent.stream_options_mut().thinking_level =
12060 Some(crate::model::ThinkingLevel::High);
12061
12062 {
12063 let cx = crate::agent_cx::AgentCx::for_request();
12064 let mut session = agent_session
12065 .session
12066 .lock(cx.cx())
12067 .await
12068 .expect("session lock");
12069 session.header.provider = Some("acme".to_string());
12070 session.header.model_id = Some("plain-model".to_string());
12071 session.header.thinking_level = None;
12072 }
12073
12074 agent_session
12075 .sync_runtime_selection_from_session_header()
12076 .await
12077 .expect("sync runtime selection");
12078
12079 assert_eq!(agent_session.agent.provider().name(), "acme");
12080 assert_eq!(agent_session.agent.provider().model_id(), "plain-model");
12081 assert_eq!(
12082 agent_session.agent.stream_options().thinking_level,
12083 Some(crate::model::ThinkingLevel::Off)
12084 );
12085
12086 let cx = crate::agent_cx::AgentCx::for_request();
12087 let session = agent_session
12088 .session
12089 .lock(cx.cx())
12090 .await
12091 .expect("session lock");
12092 assert_eq!(session.header.thinking_level.as_deref(), Some("off"));
12093 let thinking_changes = session
12094 .entries_for_current_path()
12095 .iter()
12096 .filter(|entry| {
12097 matches!(entry, crate::session::SessionEntry::ThinkingLevelChange(_))
12098 })
12099 .count();
12100 assert_eq!(thinking_changes, 1);
12101 });
12102 }
12103
12104 #[test]
12105 fn sync_runtime_selection_from_session_header_rejects_missing_credentials() {
12106 let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
12107 .build()
12108 .expect("build runtime");
12109
12110 runtime.block_on(async {
12111 let dir = tempfile::tempdir().expect("tempdir");
12112 let auth_path = dir.path().join("auth.json");
12113 let auth = AuthStorage::load(auth_path).expect("load auth");
12114 let mut agent_session = build_switch_test_session(&auth);
12115
12116 {
12117 let cx = crate::agent_cx::AgentCx::for_request();
12118 let mut session = agent_session
12119 .session
12120 .lock(cx.cx())
12121 .await
12122 .expect("session lock");
12123 session.header.provider = Some("openai".to_string());
12124 session.header.model_id = Some("gpt-4o".to_string());
12125 }
12126
12127 let err = agent_session
12128 .sync_runtime_selection_from_session_header()
12129 .await
12130 .expect_err("sync should reject switching to a credentialed target without a key");
12131 assert!(
12132 err.to_string()
12133 .contains("Missing credentials for openai/gpt-4o"),
12134 "unexpected error: {err}"
12135 );
12136 assert_eq!(agent_session.agent.provider().name(), "anthropic");
12137 assert_eq!(
12138 agent_session.agent.provider().model_id(),
12139 "claude-sonnet-4-5"
12140 );
12141
12142 let cx = crate::agent_cx::AgentCx::for_request();
12143 let session = agent_session
12144 .session
12145 .lock(cx.cx())
12146 .await
12147 .expect("session lock");
12148 assert_eq!(session.header.provider.as_deref(), Some("openai"));
12149 assert_eq!(session.header.model_id.as_deref(), Some("gpt-4o"));
12150 });
12151 }
12152
12153 #[test]
12154 fn set_provider_model_allows_current_model_without_registry() {
12155 let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
12156 .build()
12157 .expect("build runtime");
12158
12159 runtime.block_on(async {
12160 let dir = tempfile::tempdir().expect("tempdir");
12161 let auth_path = dir.path().join("auth.json");
12162 let auth = AuthStorage::load(auth_path).expect("load auth");
12163 let mut agent_session = build_switch_test_session(&auth);
12164 agent_session.model_registry = None;
12165 agent_session.agent.stream_options_mut().thinking_level =
12166 Some(crate::model::ThinkingLevel::High);
12167
12168 agent_session
12169 .set_provider_model("anthropic", "claude-sonnet-4-5")
12170 .await
12171 .expect("re-persisting the current model should succeed without a registry");
12172
12173 assert_eq!(agent_session.agent.provider().name(), "anthropic");
12174 assert_eq!(
12175 agent_session.agent.provider().model_id(),
12176 "claude-sonnet-4-5"
12177 );
12178 assert_eq!(
12179 agent_session.agent.stream_options().thinking_level,
12180 Some(crate::model::ThinkingLevel::High)
12181 );
12182
12183 let cx = crate::agent_cx::AgentCx::for_request();
12184 let session = agent_session
12185 .session
12186 .lock(cx.cx())
12187 .await
12188 .expect("session lock");
12189 assert_eq!(session.header.provider.as_deref(), Some("anthropic"));
12190 assert_eq!(
12191 session.header.model_id.as_deref(),
12192 Some("claude-sonnet-4-5")
12193 );
12194 assert_eq!(session.header.thinking_level.as_deref(), Some("high"));
12195 });
12196 }
12197
12198 #[test]
12199 fn auto_compaction_start_serializes_to_pi_mono_format() {
12200 let event = AgentEvent::AutoCompactionStart {
12201 reason: "threshold".to_string(),
12202 };
12203 let json = serde_json::to_value(&event).unwrap();
12204 assert_eq!(json["type"], "auto_compaction_start");
12205 assert_eq!(json["reason"], "threshold");
12206 }
12207
12208 #[test]
12209 fn auto_compaction_end_serializes_to_pi_mono_format() {
12210 let event = AgentEvent::AutoCompactionEnd {
12211 result: Some(serde_json::json!({
12212 "summary": "Compacted",
12213 "firstKeptEntryId": "abc123",
12214 "tokensBefore": 50000,
12215 "details": { "readFiles": [], "modifiedFiles": [] }
12216 })),
12217 aborted: false,
12218 will_retry: true,
12219 error_message: None,
12220 };
12221 let json = serde_json::to_value(&event).unwrap();
12222 assert_eq!(json["type"], "auto_compaction_end");
12223 assert!(json["result"].is_object());
12224 assert_eq!(json["aborted"], false);
12225 assert_eq!(json["willRetry"], true);
12226 assert!(json.get("errorMessage").is_none());
12227 }
12228
12229 #[test]
12230 fn auto_compaction_end_with_error_serializes_error_message() {
12231 let event = AgentEvent::AutoCompactionEnd {
12232 result: None,
12233 aborted: false,
12234 will_retry: false,
12235 error_message: Some("compaction failed".to_string()),
12236 };
12237 let json = serde_json::to_value(&event).unwrap();
12238 assert_eq!(json["type"], "auto_compaction_end");
12239 assert!(json.get("result").is_none());
12240 assert_eq!(json["errorMessage"], "compaction failed");
12241 }
12242
12243 #[test]
12244 fn apply_compaction_result_emits_structured_result_payload() {
12245 let runtime = RuntimeBuilder::current_thread()
12246 .build()
12247 .expect("runtime build");
12248
12249 runtime.block_on(async {
12250 let provider = Arc::new(SilentProvider);
12251 let tools = ToolRegistry::new(&[], Path::new("."), None);
12252 let agent = Agent::new(provider, tools, AgentConfig::default());
12253 let session = Arc::new(Mutex::new(Session::in_memory()));
12254 let agent_session =
12255 AgentSession::new(agent, session, false, ResolvedCompactionSettings::default());
12256
12257 let events: Arc<std::sync::Mutex<Vec<AgentEvent>>> =
12258 Arc::new(std::sync::Mutex::new(Vec::new()));
12259 let sink = Arc::clone(&events);
12260 let on_event: AgentEventHandler = Arc::new(move |event| {
12261 sink.lock().expect("lock compaction events").push(event);
12262 });
12263
12264 let result = compaction::CompactionResult {
12265 summary: "Compacted 10 messages into 2".to_string(),
12266 first_kept_entry_id: "entry-5".to_string(),
12267 tokens_before: 12_000,
12268 details: compaction::CompactionDetails {
12269 read_files: vec!["src/main.rs".to_string()],
12270 modified_files: vec!["src/agent.rs".to_string()],
12271 },
12272 };
12273
12274 agent_session
12275 .apply_compaction_result(result, on_event)
12276 .await
12277 .expect("apply compaction result");
12278
12279 let payload = {
12280 let guard = events.lock().expect("lock captured events");
12281 guard
12282 .iter()
12283 .find_map(|event| match event {
12284 AgentEvent::AutoCompactionEnd {
12285 result: Some(result),
12286 ..
12287 } => Some(result.clone()),
12288 _ => None,
12289 })
12290 .expect("auto compaction end payload")
12291 };
12292
12293 assert_eq!(payload["summary"], "Compacted 10 messages into 2");
12294 assert_eq!(payload["firstKeptEntryId"], "entry-5");
12295 assert_eq!(payload["tokensBefore"], 12_000);
12296 assert_eq!(payload["details"]["readFiles"], json!(["src/main.rs"]));
12297 assert_eq!(payload["details"]["modifiedFiles"], json!(["src/agent.rs"]));
12298 });
12299 }
12300
12301 #[test]
12302 fn auto_retry_start_serializes_to_pi_mono_format() {
12303 let event = AgentEvent::AutoRetryStart {
12304 attempt: 2,
12305 max_attempts: 3,
12306 delay_ms: 4000,
12307 error_message: "rate limited".to_string(),
12308 };
12309 let json = serde_json::to_value(&event).unwrap();
12310 assert_eq!(json["type"], "auto_retry_start");
12311 assert_eq!(json["attempt"], 2);
12312 assert_eq!(json["maxAttempts"], 3);
12313 assert_eq!(json["delayMs"], 4000);
12314 assert_eq!(json["errorMessage"], "rate limited");
12315 }
12316
12317 #[test]
12318 fn auto_retry_end_success_serializes_to_pi_mono_format() {
12319 let event = AgentEvent::AutoRetryEnd {
12320 success: true,
12321 attempt: 2,
12322 final_error: None,
12323 };
12324 let json = serde_json::to_value(&event).unwrap();
12325 assert_eq!(json["type"], "auto_retry_end");
12326 assert_eq!(json["success"], true);
12327 assert_eq!(json["attempt"], 2);
12328 assert!(json.get("finalError").is_none());
12329 }
12330
12331 #[test]
12332 fn auto_retry_end_failure_serializes_final_error() {
12333 let event = AgentEvent::AutoRetryEnd {
12334 success: false,
12335 attempt: 3,
12336 final_error: Some("max retries exceeded".to_string()),
12337 };
12338 let json = serde_json::to_value(&event).unwrap();
12339 assert_eq!(json["type"], "auto_retry_end");
12340 assert_eq!(json["success"], false);
12341 assert_eq!(json["attempt"], 3);
12342 assert_eq!(json["finalError"], "max retries exceeded");
12343 }
12344}