diff --git a/Cargo.lock b/Cargo.lock index eb04bbe..bb3a1a4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4648,8 +4648,6 @@ dependencies = [ [[package]] name = "mentra" version = "0.18.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94ab26ca4cd71f5c8329d5d1ed122f1de3ce75d2331702466ce0e1adc41b6c3b" dependencies = [ "async-trait", "base64", diff --git a/Cargo.toml b/Cargo.toml index 2672688..49d9d42 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -27,7 +27,7 @@ members = [ "tools/concurrency-audit", "tools/performance", ] -exclude = ["crates/libremetaverse-openjpeg", "vendor/mentra-provider"] +exclude = ["crates/libremetaverse-openjpeg", "vendor/mentra", "vendor/mentra-provider"] default-members = [ "crates/metacrate", "crates/libremetaverse-types", @@ -86,3 +86,6 @@ incremental = false [patch.crates-io] # Mentra 0.5.1 duplicates `v1` for nested OpenAI-compatible base URLs. mentra-provider = { path = "vendor/mentra-provider" } +# Mentra 0.18.3 lacks the host-facing bounded memory browse API required by +# the grid-agent operator control plane. Keep the storage schema inside Mentra. +mentra = { path = "vendor/mentra" } diff --git a/crates/metacrate-grid-agent/src/control_plane.rs b/crates/metacrate-grid-agent/src/control_plane.rs index 5778e45..80d4e79 100644 --- a/crates/metacrate-grid-agent/src/control_plane.rs +++ b/crates/metacrate-grid-agent/src/control_plane.rs @@ -25,6 +25,7 @@ const MAX_TOKEN_BYTES: usize = 16 * 1024; const MAX_OPERATOR_MESSAGE_BYTES: usize = 16 * 1024; const MAX_PAGE_SIZE: u16 = 100; const MAX_EVENT_BYTES: usize = 4 * 1024; +const MAX_MEMORY_QUERY_BYTES: usize = 1_024; #[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] @@ -248,6 +249,81 @@ pub struct AuditEventView { pub authorization_id: Option, } +#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum MemoryKindView { + Episode, + Summary, + Fact, +} + +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum MemorySortView { + #[default] + Newest, + Oldest, + Relevance, +} + +#[derive(Clone, Debug, Default, Eq, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct MemoryFilterView { + pub kind: Option, + pub pinned: Option, + pub source: Option, + pub created_from: Option, + pub created_to: Option, +} + +#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)] +pub struct MemoryCursorView { + pub created_at: i64, + pub record_id: String, +} + +#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)] +pub struct MemoryHealthView { + pub agent_count: usize, + pub record_count: usize, +} + +#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)] +pub struct MemoryAgentView { + pub avatar_id: String, + pub logical_agent_id: String, + pub mentra_agent_id: String, + pub record_count: usize, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct MemoryRecordView { + pub record_id: String, + pub kind: MemoryKindView, + pub preview: String, + pub source: Option, + pub source_revision: u64, + pub pinned: bool, + pub created_at: i64, + pub score: Option, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct MemoryPageView { + pub items: Vec, + pub next_cursor: Option, + pub next_offset: Option, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct MemoryDetailView { + pub record: MemoryRecordView, + pub content: String, + pub metadata_json: String, + pub content_truncated: bool, + pub metadata_truncated: bool, +} + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(tag = "kind", content = "data", rename_all = "snake_case")] pub enum ControlPayload { @@ -259,6 +335,10 @@ pub enum ControlPayload { PendingApprovals(Page), AuditEvents(Page), ObservabilityEvents(Page), + MemoryHealth(MemoryHealthView), + MemoryAgents(Page), + Memories(MemoryPageView), + MemoryDetail(MemoryDetailView), Accepted { operation_id: String }, Completed, Subscribed { current_sequence: u64 }, @@ -293,6 +373,32 @@ pub enum ControlRequest { ListObservabilityEvents { page: PageRequest, }, + MemoryHealth, + ListMemoryAgents { + page: PageRequest, + }, + BrowseMemories { + mentra_agent_id: String, + cursor: Option, + limit: u16, + filter: MemoryFilterView, + sort: MemorySortView, + }, + SearchMemories { + mentra_agent_id: String, + query: String, + page: PageRequest, + filter: MemoryFilterView, + sort: MemorySortView, + }, + GetMemory { + mentra_agent_id: String, + record_id: String, + }, + ForgetMemory { + mentra_agent_id: String, + record_id: String, + }, SubscribeEvents { after_sequence: Option, }, @@ -347,6 +453,11 @@ impl ControlRequest { | Self::ListPendingApprovals { .. } | Self::ListAuditEvents { .. } | Self::ListObservabilityEvents { .. } + | Self::MemoryHealth + | Self::ListMemoryAgents { .. } + | Self::BrowseMemories { .. } + | Self::SearchMemories { .. } + | Self::GetMemory { .. } | Self::SubscribeEvents { .. } ) } @@ -357,7 +468,44 @@ impl ControlRequest { | Self::ListScheduledJobs { page } | Self::ListPendingApprovals { page } | Self::ListAuditEvents { page } - | Self::ListObservabilityEvents { page } => page.valid(), + | Self::ListObservabilityEvents { page } + | Self::ListMemoryAgents { page } => page.valid(), + Self::BrowseMemories { + mentra_agent_id, + cursor, + limit, + filter, + sort, + } => { + valid_identifier(mentra_agent_id) + && (1..=MAX_PAGE_SIZE).contains(limit) + && *sort != MemorySortView::Relevance + && cursor + .as_ref() + .is_none_or(|cursor| valid_identifier(&cursor.record_id)) + && memory_filter_valid(filter) + } + Self::SearchMemories { + mentra_agent_id, + query, + page, + filter, + .. + } => { + valid_identifier(mentra_agent_id) + && query.len() <= MAX_MEMORY_QUERY_BYTES + && !query.contains('\0') + && page.valid() + && memory_filter_valid(filter) + } + Self::GetMemory { + mentra_agent_id, + record_id, + } + | Self::ForgetMemory { + mentra_agent_id, + record_id, + } => valid_identifier(mentra_agent_id) && valid_identifier(record_id), Self::CancelRequest { target_request_id } => valid_identifier(target_request_id), Self::CancelAction { action_id } => valid_identifier(action_id), Self::DecideApproval { approval_id, .. } => *approval_id != 0, @@ -1726,12 +1874,25 @@ fn valid_uuid_text(value: &str) -> bool { libremetaverse_types::UUID::parse(value.to_owned()).is_ok() } +fn memory_filter_valid(filter: &MemoryFilterView) -> bool { + filter + .source + .as_ref() + .is_none_or(|source| !source.is_empty() && source.len() <= 256 && !source.contains('\0')) + && match (filter.created_from, filter.created_to) { + (Some(from), Some(to)) => from <= to, + _ => true, + } +} + fn payload_valid(payload: &ControlPayload, maximum_frame_bytes: usize) -> bool { let page_is_bounded = match payload { ControlPayload::Sessions(page) => page.items.len() <= usize::from(MAX_PAGE_SIZE), ControlPayload::ScheduledJobs(page) => page.items.len() <= usize::from(MAX_PAGE_SIZE), ControlPayload::PendingApprovals(page) => page.items.len() <= usize::from(MAX_PAGE_SIZE), ControlPayload::AuditEvents(page) => page.items.len() <= usize::from(MAX_PAGE_SIZE), + ControlPayload::MemoryAgents(page) => page.items.len() <= usize::from(MAX_PAGE_SIZE), + ControlPayload::Memories(page) => page.items.len() <= usize::from(MAX_PAGE_SIZE), _ => true, }; page_is_bounded @@ -1771,6 +1932,12 @@ fn request_name(request: &ControlRequest) -> &'static str { ControlRequest::ListPendingApprovals { .. } => "list_pending_approvals", ControlRequest::ListAuditEvents { .. } => "list_audit_events", ControlRequest::ListObservabilityEvents { .. } => "list_observability_events", + ControlRequest::MemoryHealth => "memory_health", + ControlRequest::ListMemoryAgents { .. } => "list_memory_agents", + ControlRequest::BrowseMemories { .. } => "browse_memories", + ControlRequest::SearchMemories { .. } => "search_memories", + ControlRequest::GetMemory { .. } => "get_memory", + ControlRequest::ForgetMemory { .. } => "forget_memory", ControlRequest::SubscribeEvents { .. } => "subscribe_events", ControlRequest::CancelRequest { .. } => "cancel_request", ControlRequest::PauseAutonomy => "pause_autonomy", diff --git a/crates/metacrate-grid-agent/src/control_runtime.rs b/crates/metacrate-grid-agent/src/control_runtime.rs index 7dd50f8..bc9a25a 100644 --- a/crates/metacrate-grid-agent/src/control_runtime.rs +++ b/crates/metacrate-grid-agent/src/control_runtime.rs @@ -6,8 +6,9 @@ use crate::behavior::{BehaviorIngress, BehaviorMode}; use crate::control_plane::{ AuditEventView, CONTROL_PROTOCOL_VERSION, ControlContext, ControlError, ControlErrorCode, ControlFuture, ControlPayload, ControlRequest, ControlTarget, ConversationChannelView, - HealthView, Page, PageRequest, PendingApprovalView, RuntimeView, ScheduledJobView, - SessionMetadataView, + HealthView, MemoryAgentView, MemoryCursorView, MemoryDetailView, MemoryFilterView, + MemoryHealthView, MemoryKindView, MemoryPageView, MemoryRecordView, MemorySortView, Page, + PageRequest, PendingApprovalView, RuntimeView, ScheduledJobView, SessionMetadataView, }; use crate::conversation::{ConversationChannel, ConversationKey, ConversationStore}; use crate::observability::{ @@ -88,6 +89,7 @@ pub struct AgentControlTarget { observability: Mutex>>, build: Mutex>>, vision: Mutex>>, + memory: Mutex>>, behavior: BehaviorIngress, commands: mpsc::Sender, command_capacity: usize, @@ -133,6 +135,7 @@ impl AgentControlTarget { observability: Mutex::new(None), build: Mutex::new(None), vision: Mutex::new(None), + memory: Mutex::new(None), behavior, commands, command_capacity, @@ -158,6 +161,10 @@ impl AgentControlTarget { *lock(&self.vision) = Some(vision); } + pub fn attach_memory_control(&self, memory: Arc) { + *lock(&self.memory) = Some(memory); + } + pub fn update_session(&self, session: SessionStatus) { let mut state = lock(&self.state); state.session = session; @@ -362,6 +369,156 @@ impl AgentControlTarget { &page, values, ))) } + ControlRequest::MemoryHealth => { + let memory = memory_admin(&self.memory)?; + let (agent_count, record_count) = memory.total_counts().map_err(memory_error)?; + Ok(ControlPayload::MemoryHealth(MemoryHealthView { + agent_count, + record_count, + })) + } + ControlRequest::ListMemoryAgents { page } => { + require_operator(context)?; + let values = memory_admin(&self.memory)? + .agents() + .map_err(memory_error)? + .into_iter() + .map(|(mapping, record_count)| MemoryAgentView { + avatar_id: mapping.avatar_id, + logical_agent_id: mapping.logical_agent_id, + mentra_agent_id: mapping.mentra_agent_id, + record_count, + }) + .collect(); + Ok(ControlPayload::MemoryAgents(page_values(&page, values))) + } + ControlRequest::BrowseMemories { + mentra_agent_id, + cursor, + limit, + filter, + sort, + } => { + require_operator(context)?; + let (records, next_cursor) = memory_admin(&self.memory)? + .browse( + &mentra_agent_id, + cursor.map(|cursor| mentra::memory::MemoryListCursor { + created_at: cursor.created_at, + record_id: cursor.record_id, + }), + usize::from(limit), + memory_filter(filter), + memory_sort(sort)?, + ) + .map_err(memory_error)?; + Ok(ControlPayload::Memories(MemoryPageView { + items: records.into_iter().map(memory_record_view).collect(), + next_cursor: next_cursor.map(|cursor| MemoryCursorView { + created_at: cursor.created_at, + record_id: cursor.record_id, + }), + next_offset: None, + })) + } + ControlRequest::SearchMemories { + mentra_agent_id, + query, + page, + filter, + sort, + } => { + require_operator(context)?; + let start = usize::try_from(page.cursor.unwrap_or(0)).unwrap_or(usize::MAX); + let requested = start + .saturating_add(usize::from(page.limit)) + .saturating_add(1); + let memory = memory_admin(&self.memory)?; + let mut records = if query.trim().is_empty() { + memory + .browse( + &mentra_agent_id, + None, + requested.min(500), + memory_filter(filter.clone()), + match sort { + MemorySortView::Oldest => mentra::memory::MemoryListSort::Oldest, + _ => mentra::memory::MemoryListSort::Newest, + }, + ) + .map(|(records, _)| records) + } else { + memory.search( + &mentra_agent_id, + &query, + requested.min(500), + memory_filter(filter.clone()), + ) + } + .map_err(memory_error)?; + records.retain(|record| record_matches(record, &filter)); + match sort { + MemorySortView::Newest => records.sort_by(|a, b| { + (b.created_at, &b.record_id).cmp(&(a.created_at, &a.record_id)) + }), + MemorySortView::Oldest => records.sort_by(|a, b| { + (a.created_at, &a.record_id).cmp(&(b.created_at, &b.record_id)) + }), + MemorySortView::Relevance => {} + } + let has_more = records.len() > start.saturating_add(usize::from(page.limit)); + let items = records + .into_iter() + .skip(start) + .take(usize::from(page.limit)) + .map(memory_record_view) + .collect(); + Ok(ControlPayload::Memories(MemoryPageView { + items, + next_cursor: None, + next_offset: has_more.then(|| { + u64::try_from(start.saturating_add(usize::from(page.limit))) + .unwrap_or(u64::MAX) + }), + })) + } + ControlRequest::GetMemory { + mentra_agent_id, + record_id, + } => { + require_operator(context)?; + let record = memory_admin(&self.memory)? + .detail(&mentra_agent_id, &record_id) + .map_err(memory_error)? + .ok_or_else(|| { + control_error(ControlErrorCode::NotFound, "memory record not found", false) + })?; + Ok(ControlPayload::MemoryDetail(MemoryDetailView { + record: memory_record_view(record.clone()), + content: truncate_utf8(&record.content, 24 * 1024), + metadata_json: truncate_utf8(&record.metadata_json, 8 * 1024), + content_truncated: record.content.len() > 24 * 1024, + metadata_truncated: record.metadata_json.len() > 8 * 1024, + })) + } + ControlRequest::ForgetMemory { + mentra_agent_id, + record_id, + } => { + require_operator(context)?; + if memory_admin(&self.memory)? + .forget(&mentra_agent_id, &record_id) + .map_err(memory_error)? + { + Ok(ControlPayload::Completed) + } else { + Err(control_error( + ControlErrorCode::NotFound, + "memory record not found", + false, + )) + } + } ControlRequest::PauseAutonomy => { let response = self.enqueue("pause", RuntimeControlCommand::Pause)?; self.behavior.pause(); @@ -639,6 +796,25 @@ fn runtime_domain_event( .result_code("completed")? .code_field("decision", if *approve { "approved" } else { "denied" })?, ), + ControlRequest::ForgetMemory { + mentra_agent_id, + record_id, + } => Some( + EventDraft::new( + EventFamily::ControlCommand, + EventSeverity::Info, + "memory", + EventOrigin::Control, + )? + .correlation(CorrelationIds { + session_id: Some(pseudonymous_identifier(mentra_agent_id)), + action_id: Some(pseudonymous_identifier(record_id)), + ..CorrelationIds::default() + })? + .result_code("completed")? + .code_field("operation", "forget_memory")? + .redacted("memory_content")?, + ), ControlRequest::GracefulShutdown => Some( EventDraft::new( EventFamily::Shutdown, @@ -659,6 +835,7 @@ fn request_correlation(request: &ControlRequest) -> CorrelationIds { ControlRequest::DecideApproval { approval_id, .. } => { Some(format!("approval-{approval_id}")) } + ControlRequest::ForgetMemory { record_id, .. } => Some(record_id.clone()), _ => None, }; CorrelationIds { @@ -677,6 +854,12 @@ const fn runtime_request_name(request: &ControlRequest) -> &'static str { ControlRequest::ListPendingApprovals { .. } => "list_pending_approvals", ControlRequest::ListAuditEvents { .. } => "list_audit_events", ControlRequest::ListObservabilityEvents { .. } => "list_observability_events", + ControlRequest::MemoryHealth => "memory_health", + ControlRequest::ListMemoryAgents { .. } => "list_memory_agents", + ControlRequest::BrowseMemories { .. } => "browse_memories", + ControlRequest::SearchMemories { .. } => "search_memories", + ControlRequest::GetMemory { .. } => "get_memory", + ControlRequest::ForgetMemory { .. } => "forget_memory", ControlRequest::SubscribeEvents { .. } => "subscribe_events", ControlRequest::CancelRequest { .. } => "cancel_request", ControlRequest::PauseAutonomy => "pause_autonomy", @@ -753,6 +936,111 @@ const fn policy_outcome_name(outcome: PolicyFinalOutcome) -> &'static str { } } +fn memory_admin( + slot: &Mutex>>, +) -> Result, ControlError> { + lock(slot).clone().ok_or_else(|| { + control_error( + ControlErrorCode::Conflict, + "memory service is not available", + true, + ) + }) +} + +fn require_operator(context: &ControlContext) -> Result<(), ControlError> { + (context.role == crate::control_plane::ControlRole::Operator) + .then_some(()) + .ok_or_else(|| { + control_error( + ControlErrorCode::PermissionDenied, + "memory content requires operator role", + false, + ) + }) +} + +fn memory_filter(filter: MemoryFilterView) -> mentra::memory::MemoryListFilter { + mentra::memory::MemoryListFilter { + kind: filter.kind.map(|kind| match kind { + MemoryKindView::Episode => mentra::memory::MemoryRecordKind::Episode, + MemoryKindView::Summary => mentra::memory::MemoryRecordKind::Summary, + MemoryKindView::Fact => mentra::memory::MemoryRecordKind::Fact, + }), + pinned: filter.pinned, + source: filter.source, + created_from: filter.created_from, + created_to: filter.created_to, + } +} + +fn memory_sort(sort: MemorySortView) -> Result { + match sort { + MemorySortView::Newest => Ok(mentra::memory::MemoryListSort::Newest), + MemorySortView::Oldest => Ok(mentra::memory::MemoryListSort::Oldest), + MemorySortView::Relevance => Err(control_error( + ControlErrorCode::InvalidRequest, + "relevance sort requires a search query", + false, + )), + } +} + +fn record_matches(record: &mentra::memory::MemoryRecord, filter: &MemoryFilterView) -> bool { + filter.kind.is_none_or(|kind| { + matches!( + (kind, record.kind), + ( + MemoryKindView::Episode, + mentra::memory::MemoryRecordKind::Episode + ) | ( + MemoryKindView::Summary, + mentra::memory::MemoryRecordKind::Summary + ) | (MemoryKindView::Fact, mentra::memory::MemoryRecordKind::Fact) + ) + }) && filter.pinned.is_none_or(|pinned| record.pinned == pinned) + && filter + .source + .as_ref() + .is_none_or(|source| record.source.as_ref() == Some(source)) + && filter + .created_from + .is_none_or(|from| record.created_at >= from) + && filter.created_to.is_none_or(|to| record.created_at <= to) +} + +fn memory_record_view(record: mentra::memory::MemoryRecord) -> MemoryRecordView { + MemoryRecordView { + record_id: record.record_id, + kind: match record.kind { + mentra::memory::MemoryRecordKind::Episode => MemoryKindView::Episode, + mentra::memory::MemoryRecordKind::Summary => MemoryKindView::Summary, + mentra::memory::MemoryRecordKind::Fact => MemoryKindView::Fact, + }, + preview: truncate_utf8(&record.content, 160), + source: record.source.map(|source| truncate_utf8(&source, 128)), + source_revision: record.source_revision, + pinned: record.pinned, + created_at: record.created_at, + score: record.score, + } +} + +fn truncate_utf8(value: &str, maximum: usize) -> String { + if value.len() <= maximum { + return value.to_owned(); + } + let mut end = maximum; + while !value.is_char_boundary(end) { + end -= 1; + } + value[..end].to_owned() +} + +fn memory_error(_: String) -> ControlError { + control_error(ControlErrorCode::Internal, "memory operation failed", false) +} + fn unix_seconds() -> u64 { SystemTime::now() .duration_since(UNIX_EPOCH) diff --git a/crates/metacrate-grid-agent/src/control_runtime_tests.rs b/crates/metacrate-grid-agent/src/control_runtime_tests.rs index f8fb9bb..77af690 100644 --- a/crates/metacrate-grid-agent/src/control_runtime_tests.rs +++ b/crates/metacrate-grid-agent/src/control_runtime_tests.rs @@ -1,7 +1,7 @@ use crate::behavior::{BehaviorController, BehaviorError, EmbodimentFuture, EmbodimentSink}; use crate::control_plane::{ CONTROL_PROTOCOL_VERSION, ControlLimits, ControlPayload, ControlPlane, ControlRequest, - ControlRequestEnvelope, PageRequest, + ControlRequestEnvelope, MemoryFilterView, MemorySortView, PageRequest, }; use crate::control_runtime::{AgentControlTarget, RuntimeControlCommand}; use crate::conversation::{ConversationLimits, ConversationStore}; @@ -12,6 +12,7 @@ use crate::{ }; use libremetaverse_types::UUID; use libremetaverse_types::compat::CancellationToken; +use mentra::memory::{MemoryRecord, MemoryRecordKind, MemoryStore}; use std::collections::BTreeSet; use std::sync::{Arc, Mutex}; use std::time::Duration; @@ -152,6 +153,31 @@ async fn production_target_projects_state_and_routes_real_mutations() { target.attach_observability(observability); let vision = Arc::new(FakeVision(Mutex::new(Some("visual-1".into())))); target.attach_vision_control(vision.clone()); + let memory_root = + std::env::temp_dir().join(format!("metacrate-control-memory-{}", std::process::id())); + let memory_store = mentra::runtime::HybridRuntimeStore::new(memory_root.join("runtime.sqlite")); + memory_store + .upsert_records(&[MemoryRecord { + record_id: "fact:operator".into(), + agent_id: "mentra-agent".into(), + kind: MemoryRecordKind::Fact, + content: format!("operator-only durable detail {}", "x".repeat(100_000)), + source_revision: 1, + created_at: 1, + metadata_json: "{}".into(), + source: Some("manual".into()), + pinned: true, + score: None, + }]) + .expect("seed memory"); + let memory = Arc::new( + crate::MentraMemoryAdmin::new(memory_store, memory_root.join("memory-agents.json")) + .expect("memory admin"), + ); + memory + .register("avatar-id", "logical-agent", "mentra-agent") + .expect("mapping"); + target.attach_memory_control(memory); target.update_region( Some("00000000-0000-4000-8000-000000000001".into()), Some(9_007_199_254_740_992), @@ -200,6 +226,81 @@ async fn production_target_projects_state_and_routes_real_mutations() { assert!(vision.active_capture().is_none(), "pause cancels capture"); let observer = plane.connect("observer-token").expect("observer"); + let health = observer + .request(envelope("memory-health", ControlRequest::MemoryHealth)) + .await; + assert!(matches!(health.result, Ok(ControlPayload::MemoryHealth(_)))); + let denied_memory = observer + .request(envelope( + "memory-denied", + ControlRequest::BrowseMemories { + mentra_agent_id: "mentra-agent".into(), + cursor: None, + limit: 10, + filter: MemoryFilterView::default(), + sort: MemorySortView::Newest, + }, + )) + .await; + assert!(denied_memory.result.is_err()); + let memories = operator + .request(envelope( + "memory-browse", + ControlRequest::BrowseMemories { + mentra_agent_id: "mentra-agent".into(), + cursor: None, + limit: 10, + filter: MemoryFilterView::default(), + sort: MemorySortView::Newest, + }, + )) + .await; + let Ok(ControlPayload::Memories(memories)) = memories.result else { + panic!("operator memory browse") + }; + assert_eq!(memories.items.len(), 1); + assert!(memories.items[0].preview.contains("durable detail")); + for (id, query) in [("memory-empty", ""), ("memory-punctuation", "(durable)!!!")] { + let search = operator + .request(envelope( + id, + ControlRequest::SearchMemories { + mentra_agent_id: "mentra-agent".into(), + query: query.into(), + page: PageRequest { + cursor: None, + limit: 10, + }, + filter: MemoryFilterView { + source: Some("manual".into()), + ..MemoryFilterView::default() + }, + sort: if query.is_empty() { + MemorySortView::Newest + } else { + MemorySortView::Relevance + }, + }, + )) + .await; + let Ok(ControlPayload::Memories(search)) = search.result else { + panic!("memory search") + }; + assert_eq!(search.items.len(), 1); + } + let detail = operator + .request(envelope( + "memory-detail", + ControlRequest::GetMemory { + mentra_agent_id: "mentra-agent".into(), + record_id: "fact:operator".into(), + }, + )) + .await; + let Ok(ControlPayload::MemoryDetail(detail)) = detail.result else { + panic!("bounded memory detail") + }; + assert_eq!(detail.content.len(), 24 * 1024); let denied = observer .request(envelope("shutdown", ControlRequest::GracefulShutdown)) .await; diff --git a/crates/metacrate-grid-agent/src/interaction.rs b/crates/metacrate-grid-agent/src/interaction.rs index 6ff2657..f1f78d3 100644 --- a/crates/metacrate-grid-agent/src/interaction.rs +++ b/crates/metacrate-grid-agent/src/interaction.rs @@ -1779,7 +1779,8 @@ pub struct PolicyLlmResponder { client: Arc, approval_reviewer: Arc, runtime: Arc, - agents: Mutex>>>, + agents: Mutex>, + memory_admin: Arc, active_tools: Arc>>, gateway: Arc, backend: Arc, @@ -1788,6 +1789,12 @@ pub struct PolicyLlmResponder { storage_path: std::path::PathBuf, } +#[derive(Clone)] +struct MentraAgentHandle { + id: String, + agent: Arc>, +} + #[derive(Clone)] struct ActiveMentraRequest { executor: Arc, @@ -1994,11 +2001,17 @@ impl PolicyLlmResponder { let active_tools = Arc::new(Mutex::new(BTreeMap::new())); std::fs::create_dir_all(&storage_path) .map_err(|error| InteractionError::Persistence(error.to_string()))?; + let store = mentra::runtime::HybridRuntimeStore::new(storage_path.join("runtime.sqlite")); + let memory_admin = Arc::new( + crate::memory_admin::MentraMemoryAdmin::new( + store.clone(), + storage_path.join("memory-agents.json"), + ) + .map_err(InteractionError::Persistence)?, + ); let mut builder = mentra::Runtime::builder() .with_runtime_identifier(MENTRA_RUNTIME_IDENTIFIER) - .with_store(mentra::runtime::HybridRuntimeStore::new( - storage_path.join("runtime.sqlite"), - )) + .with_store(store) .with_registered_provider(client.mentra_provider()); for definition in gateway.registered_tool_definitions() { builder = builder.with_tool(MentraGridTool { @@ -2015,9 +2028,13 @@ impl PolicyLlmResponder { .into_iter() .filter(|agent| agent.name().starts_with(MENTRA_AGENT_PREFIX)) .map(|agent| { + let id = agent.id().to_owned(); ( agent.name().to_owned(), - Arc::new(tokio::sync::Mutex::new(agent)), + MentraAgentHandle { + id, + agent: Arc::new(tokio::sync::Mutex::new(agent)), + }, ) }) .collect(); @@ -2029,6 +2046,7 @@ impl PolicyLlmResponder { client, runtime: Arc::new(runtime), agents: Mutex::new(agents), + memory_admin, active_tools, gateway, backend, @@ -2048,8 +2066,13 @@ impl PolicyLlmResponder { .agents .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - if let Some(agent) = agents.get(&key) { - return Ok(Arc::clone(agent)); + if let Some(handle) = agents.get(&key) { + if request.origin == InteractionOrigin::AuthorizedIm { + self.memory_admin + .register(&request.sender_id.to_string(), &key, &handle.id) + .map_err(|_| InteractionModelError::Failed)?; + } + return Ok(Arc::clone(&handle.agent)); } let tools = self.gateway.tools_for(context, (self.now)()); let mut tool_names = tools @@ -2081,10 +2104,27 @@ impl PolicyLlmResponder { .runtime .spawn_with_config(key.clone(), model, config) .map_err(|_| InteractionModelError::Failed)?; + let id = agent.id().to_owned(); let agent = Arc::new(tokio::sync::Mutex::new(agent)); - agents.insert(key, Arc::clone(&agent)); + agents.insert( + key.clone(), + MentraAgentHandle { + id: id.clone(), + agent: Arc::clone(&agent), + }, + ); + if request.origin == InteractionOrigin::AuthorizedIm { + self.memory_admin + .register(&request.sender_id.to_string(), &key, &id) + .map_err(|_| InteractionModelError::Failed)?; + } Ok(agent) } + + #[must_use] + pub fn memory_admin(&self) -> Arc { + Arc::clone(&self.memory_admin) + } } pub(crate) fn mentra_agent_name(request: &ResponseRequest) -> String { diff --git a/crates/metacrate-grid-agent/src/lib.rs b/crates/metacrate-grid-agent/src/lib.rs index 3fc9cdd..13a558b 100644 --- a/crates/metacrate-grid-agent/src/lib.rs +++ b/crates/metacrate-grid-agent/src/lib.rs @@ -15,6 +15,7 @@ pub mod conversation; pub mod interaction; pub mod landmarks; pub mod llm; +pub mod memory_admin; pub mod observability; pub mod perception; pub mod policy; @@ -96,8 +97,10 @@ pub use control_plane::{ ControlErrorCode, ControlEvent, ControlEventKind, ControlFuture, ControlLimits, ControlPayload, ControlPlane, ControlRequest, ControlRequestEnvelope, ControlResponseEnvelope, ControlRole, ControlSubscription, ControlTarget, ConversationChannelView, HealthView, - InProcessControlClient, Page, PageRequest, PendingApprovalView, RuntimeView, ScheduledJobView, - SessionMetadataView, TcpControlClient, TcpControlConfig, TcpControlServer, + InProcessControlClient, MemoryAgentView, MemoryCursorView, MemoryDetailView, MemoryFilterView, + MemoryHealthView, MemoryKindView, MemoryPageView, MemoryRecordView, MemorySortView, Page, + PageRequest, PendingApprovalView, RuntimeView, ScheduledJobView, SessionMetadataView, + TcpControlClient, TcpControlConfig, TcpControlServer, }; pub use control_runtime::{AgentControlTarget, RuntimeControlCommand}; pub use conversation::{ @@ -130,6 +133,7 @@ pub use landmarks::{ pub use llm::{ CompletionMessage, ContentPart, ImageDetail, LlmClient, LlmError, ToolDefinition, ToolSchema, }; +pub use memory_admin::{MemoryAgentMapping, MentraMemoryAdmin}; pub use observability::{ AgentMetrics, CorrelationIds, DiagnosticDirection, EventDraft, EventFamily, EventOrigin, EventSeverity, EventSubscription, JournalConfig, LatencyMetrics, MetricsSnapshot, diff --git a/crates/metacrate-grid-agent/src/main.rs b/crates/metacrate-grid-agent/src/main.rs index 9dd1679..734cda7 100644 --- a/crates/metacrate-grid-agent/src/main.rs +++ b/crates/metacrate-grid-agent/src/main.rs @@ -415,6 +415,7 @@ async fn run_live( control_target.attach_observability(live.observability.clone()); control_target.attach_build_control(live.build_control.clone()); control_target.attach_vision_control(live.vision.clone()); + control_target.attach_memory_control(live.memory_admin.clone()); control_target.update_session(handle.status()); let erased_target: Arc = control_target.clone(); let (control_plane, integrated_client, control_server) = match config.mode { @@ -732,6 +733,7 @@ struct LiveInteractions { policy: Arc, audit: Arc, observability: Arc, + memory_admin: Arc, _appearance_recovery: libremetaverse_types::compat::Subscription, _landmark_intake: metacrate_grid_agent::LibremetaverseLandmarkIntake, landmark_roaming: metacrate_grid_agent::LandmarkRoamingHandle, @@ -903,6 +905,7 @@ async fn start_live_interactions( now, config.storage_path.join("mentra"), )?); + let memory_admin = responder.memory_admin(); let vision_limits = config.vision; let scene_source = Arc::new(LibremetaverseSceneSource::new(owner, vision_limits)); scene_source.start_prefetch(); @@ -936,6 +939,7 @@ async fn start_live_interactions( policy: gateway, audit, observability, + memory_admin, _appearance_recovery: appearance_recovery, _landmark_intake: landmark_intake, landmark_roaming, diff --git a/crates/metacrate-grid-agent/src/memory_admin.rs b/crates/metacrate-grid-agent/src/memory_admin.rs new file mode 100644 index 0000000..29a6378 --- /dev/null +++ b/crates/metacrate-grid-agent/src/memory_admin.rs @@ -0,0 +1,288 @@ +//! Operator-only access to durable Mentra memory. + +#![allow(clippy::missing_errors_doc)] + +use mentra::memory::{ + MemoryListCursor, MemoryListFilter, MemoryListRequest, MemoryListSort, MemoryRecord, + MemorySearchMode, MemorySearchRequest, MemoryStore, +}; +use mentra::runtime::HybridRuntimeStore; +use serde::{Deserialize, Serialize}; +use std::collections::BTreeMap; +use std::path::{Path, PathBuf}; +use std::sync::{Mutex, MutexGuard}; + +const MAX_MAPPINGS: usize = 4_096; + +#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)] +pub struct MemoryAgentMapping { + pub avatar_id: String, + pub logical_agent_id: String, + pub mentra_agent_id: String, +} + +pub struct MentraMemoryAdmin { + store: HybridRuntimeStore, + mappings_path: PathBuf, + mappings: Mutex>, +} + +impl MentraMemoryAdmin { + pub fn new(store: HybridRuntimeStore, mappings_path: PathBuf) -> Result { + let mappings = load_mappings(&mappings_path)?; + Ok(Self { + store, + mappings_path, + mappings: Mutex::new(mappings), + }) + } + + pub fn register( + &self, + avatar_id: &str, + logical_agent_id: &str, + mentra_agent_id: &str, + ) -> Result<(), String> { + let mapping = MemoryAgentMapping { + avatar_id: avatar_id.to_owned(), + logical_agent_id: logical_agent_id.to_owned(), + mentra_agent_id: mentra_agent_id.to_owned(), + }; + if !valid_mapping(&mapping) { + return Err("invalid memory-agent mapping".to_owned()); + } + let mut mappings = lock(&self.mappings); + if !mappings.contains_key(avatar_id) && mappings.len() >= MAX_MAPPINGS { + return Err("memory-agent mapping limit reached".to_owned()); + } + if mappings.get(avatar_id).is_some_and(|mapping| { + mapping.logical_agent_id == logical_agent_id + && mapping.mentra_agent_id == mentra_agent_id + }) { + return Ok(()); + } + let previous = mappings.insert(avatar_id.to_owned(), mapping); + if let Err(error) = persist_mappings(&self.mappings_path, mappings.values()) { + if let Some(previous) = previous { + mappings.insert(avatar_id.to_owned(), previous); + } else { + mappings.remove(avatar_id); + } + return Err(error); + } + Ok(()) + } + + pub fn agents(&self) -> Result, String> { + lock(&self.mappings) + .values() + .map(|mapping| { + self.store + .count_records(&mapping.mentra_agent_id) + .map(|count| (mapping.clone(), count)) + .map_err(|error| error.to_string()) + }) + .collect() + } + + pub fn total_counts(&self) -> Result<(usize, usize), String> { + let agents = self.agents()?; + Ok((agents.len(), agents.iter().map(|(_, count)| count).sum())) + } + + pub fn browse( + &self, + mentra_agent_id: &str, + cursor: Option, + limit: usize, + filter: MemoryListFilter, + sort: MemoryListSort, + ) -> Result<(Vec, Option), String> { + self.require_agent(mentra_agent_id)?; + self.store + .list_records(&MemoryListRequest { + agent_id: mentra_agent_id.to_owned(), + cursor, + limit, + filter, + sort, + }) + .map(|page| (page.records, page.next_cursor)) + .map_err(|error| error.to_string()) + } + + pub fn search( + &self, + mentra_agent_id: &str, + query: &str, + limit: usize, + filter: MemoryListFilter, + ) -> Result, String> { + self.require_agent(mentra_agent_id)?; + self.store + .search_records_with_options(&MemorySearchRequest { + agent_id: mentra_agent_id.to_owned(), + query: query.to_owned(), + limit, + char_budget: None, + mode: MemorySearchMode::Tool, + filter, + }) + .map_err(|error| error.to_string()) + } + + pub fn detail( + &self, + mentra_agent_id: &str, + record_id: &str, + ) -> Result, String> { + self.require_agent(mentra_agent_id)?; + self.store + .get_record(mentra_agent_id, record_id) + .map_err(|error| error.to_string()) + } + + pub fn forget(&self, mentra_agent_id: &str, record_id: &str) -> Result { + self.require_agent(mentra_agent_id)?; + self.store + .tombstone_records(mentra_agent_id, &[record_id.to_owned()]) + .map(|affected| affected == 1) + .map_err(|error| error.to_string()) + } + + fn require_agent(&self, mentra_agent_id: &str) -> Result<(), String> { + lock(&self.mappings) + .values() + .any(|mapping| mapping.mentra_agent_id == mentra_agent_id) + .then_some(()) + .ok_or_else(|| "memory agent not found".to_owned()) + } +} + +fn load_mappings(path: &Path) -> Result, String> { + let bytes = match std::fs::read(path) { + Ok(bytes) => bytes, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + match std::fs::read(path.with_extension("json.bak")) { + Ok(bytes) => bytes, + Err(backup) if backup.kind() == std::io::ErrorKind::NotFound => { + return Ok(BTreeMap::new()); + } + Err(backup) => return Err(backup.to_string()), + } + } + Err(error) => return Err(error.to_string()), + }; + let values: Vec = + serde_json::from_slice(&bytes).map_err(|e| e.to_string())?; + if values.len() > MAX_MAPPINGS || values.iter().any(|mapping| !valid_mapping(mapping)) { + return Err("invalid or excessive memory-agent mappings".to_owned()); + } + let mut mappings = BTreeMap::new(); + for mapping in values { + if mappings + .insert(mapping.avatar_id.clone(), mapping) + .is_some() + { + return Err("duplicate memory-agent mapping".to_owned()); + } + } + Ok(mappings) +} + +fn valid_mapping(mapping: &MemoryAgentMapping) -> bool { + [ + mapping.avatar_id.as_str(), + mapping.logical_agent_id.as_str(), + mapping.mentra_agent_id.as_str(), + ] + .iter() + .all(|value| !value.is_empty() && value.len() <= 256 && !value.contains('\0')) +} + +fn persist_mappings<'a>( + path: &Path, + mappings: impl Iterator, +) -> Result<(), String> { + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent).map_err(|error| error.to_string())?; + } + let bytes = serde_json::to_vec_pretty(&mappings.collect::>()) + .map_err(|error| error.to_string())?; + let temporary = path.with_extension("json.tmp"); + let backup = path.with_extension("json.bak"); + std::fs::write(&temporary, bytes).map_err(|error| error.to_string())?; + let had_previous = path.exists(); + if had_previous { + let _ = std::fs::remove_file(&backup); + std::fs::rename(path, &backup).map_err(|error| error.to_string())?; + } + if let Err(error) = std::fs::rename(&temporary, path) { + if had_previous { + let _ = std::fs::rename(&backup, path); + } + return Err(error.to_string()); + } + if had_previous { + let _ = std::fs::remove_file(backup); + } + Ok(()) +} + +fn lock(value: &Mutex) -> MutexGuard<'_, T> { + value + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) +} + +#[cfg(test)] +mod tests { + use super::*; + use mentra::memory::{MemoryRecordKind, MemoryStore}; + use std::time::{SystemTime, UNIX_EPOCH}; + + #[test] + fn mapping_persists_and_cross_agent_forget_fails_closed() { + let root = std::env::temp_dir().join(format!( + "metacrate-memory-admin-{}-{}", + std::process::id(), + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos() + )); + let database = root.join("runtime.sqlite"); + let mappings = root.join("memory-agents.json"); + let store = HybridRuntimeStore::new(&database); + store + .upsert_records(&[record("fact:a", "mentra-a"), record("fact:b", "mentra-b")]) + .unwrap(); + let admin = MentraMemoryAdmin::new(store.clone(), mappings.clone()).unwrap(); + admin.register("avatar-a", "logical-a", "mentra-a").unwrap(); + admin.register("avatar-b", "logical-b", "mentra-b").unwrap(); + assert!(!admin.forget("mentra-a", "fact:b").unwrap()); + assert!(admin.detail("mentra-b", "fact:b").unwrap().is_some()); + drop(admin); + + let reopened = MentraMemoryAdmin::new(store, mappings).unwrap(); + assert_eq!(reopened.agents().unwrap().len(), 2); + assert!(reopened.forget("mentra-a", "fact:a").unwrap()); + assert!(reopened.detail("mentra-a", "fact:a").unwrap().is_none()); + let _ = std::fs::remove_dir_all(root); + } + + fn record(id: &str, agent_id: &str) -> MemoryRecord { + MemoryRecord { + record_id: id.to_owned(), + agent_id: agent_id.to_owned(), + kind: MemoryRecordKind::Fact, + content: "durable memory".to_owned(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_owned(), + source: Some("manual".to_owned()), + pinned: true, + score: None, + } + } +} diff --git a/crates/metacrate-grid-agent/src/tui.rs b/crates/metacrate-grid-agent/src/tui.rs index 8b832ff..6c876a3 100644 --- a/crates/metacrate-grid-agent/src/tui.rs +++ b/crates/metacrate-grid-agent/src/tui.rs @@ -8,8 +8,9 @@ use crate::control_plane::{ AuditEventView, CONTROL_PROTOCOL_VERSION, ControlEvent, ControlEventKind, ControlPayload, ControlRequest, ControlRequestEnvelope, ControlResponseEnvelope, HealthView, - InProcessControlClient, Page, PageRequest, PendingApprovalView, RuntimeView, ScheduledJobView, - SessionMetadataView, TcpControlClient, + InProcessControlClient, MemoryAgentView, MemoryCursorView, MemoryDetailView, MemoryFilterView, + MemoryHealthView, MemoryKindView, MemoryRecordView, MemorySortView, Page, PageRequest, + PendingApprovalView, RuntimeView, ScheduledJobView, SessionMetadataView, TcpControlClient, }; use crate::observability::{EventSeverity, MetricsSnapshot, StructuredEvent}; use crate::{AgentPreferences, ControlLimits, OperatingMode, PreferencesSummary, SecretString}; @@ -42,6 +43,7 @@ const PAGE_LIMIT: u16 = 100; pub enum OperatorScreen { Overview, Sessions, + Memories, QueuesAndBudgets, Roaming, Approvals, @@ -53,9 +55,10 @@ pub enum OperatorScreen { } impl OperatorScreen { - const ALL: [Self; 10] = [ + const ALL: [Self; 11] = [ Self::Overview, Self::Sessions, + Self::Memories, Self::QueuesAndBudgets, Self::Roaming, Self::Approvals, @@ -71,6 +74,7 @@ impl OperatorScreen { match self { Self::Overview => "Overview", Self::Sessions => "Sessions", + Self::Memories => "Memory", Self::QueuesAndBudgets => "Queues & budgets", Self::Roaming => "Roaming", Self::Approvals => "Approvals", @@ -337,6 +341,10 @@ pub struct OperatorSnapshot { pub sessions: Vec, pub schedules: Vec, pub approvals: Vec, + pub memory_health: Option, + pub memory_agents: Vec, + pub memories: Vec, + pub memory_detail: Option, pub audit: Vec, pub timeline: VecDeque, pub gap_count: u64, @@ -352,6 +360,10 @@ impl Default for OperatorSnapshot { sessions: Vec::new(), schedules: Vec::new(), approvals: Vec::new(), + memory_health: None, + memory_agents: Vec::new(), + memories: Vec::new(), + memory_detail: None, audit: Vec::new(), timeline: VecDeque::with_capacity(MAX_TIMELINE), gap_count: 0, @@ -376,10 +388,23 @@ pub enum OperatorCommand { Pause, Resume, CancelAction(String), - DecideApproval { id: u64, approve: bool }, + DecideApproval { + id: u64, + approve: bool, + }, Reconnect, - ExpireSession { avatar_id: String, direct_im: bool }, - ToggleSchedule { job_id: String, enabled: bool }, + ExpireSession { + avatar_id: String, + direct_im: bool, + }, + ToggleSchedule { + job_id: String, + enabled: bool, + }, + ForgetMemory { + mentra_agent_id: String, + record_id: String, + }, Shutdown, } @@ -411,6 +436,13 @@ impl OperatorCommand { job_id: job_id.clone(), enabled: *enabled, }, + Self::ForgetMemory { + mentra_agent_id, + record_id, + } => ControlRequest::ForgetMemory { + mentra_agent_id: mentra_agent_id.clone(), + record_id: record_id.clone(), + }, Self::Shutdown => ControlRequest::GracefulShutdown, } } @@ -424,7 +456,7 @@ impl OperatorCommand { | Self::Reconnect | Self::ExpireSession { .. } | Self::ToggleSchedule { .. } => CommandConfirmation::Normal, - Self::Shutdown => CommandConfirmation::HighRisk, + Self::ForgetMemory { .. } | Self::Shutdown => CommandConfirmation::HighRisk, } } } @@ -459,13 +491,21 @@ pub enum TuiInput { FilterBackspace, FilterCommit, FilterCancel, + MemoryLoad, + MemoryNextAgent, + MemoryNextPage, + MemoryPreviousPage, + MemoryCycleKind, + MemoryTogglePinned, + MemoryCycleSort, } -#[derive(Clone, Debug, Eq, PartialEq)] +#[derive(Clone, Debug, PartialEq)] pub enum TuiAction { None, Refresh, Execute(OperatorCommand), + Request(ControlRequest), Exit, } @@ -605,6 +645,14 @@ pub struct OperatorTui { cadence_millis: ValueCell, last_frame: ValueCell>, filter_editing: bool, + memory_agent: usize, + memory_filter: MemoryFilterView, + memory_sort: MemorySortView, + memory_next_cursor: Option, + memory_cursor_stack: Vec>, + memory_current_offset: u64, + memory_next_offset: Option, + memory_refresh_pending: bool, } impl Default for OperatorTui { @@ -627,6 +675,14 @@ impl Default for OperatorTui { cadence_millis: ValueCell::new(0), last_frame: ValueCell::new(None), filter_editing: false, + memory_agent: 0, + memory_filter: MemoryFilterView::default(), + memory_sort: MemorySortView::Newest, + memory_next_cursor: None, + memory_cursor_stack: Vec::new(), + memory_current_offset: 0, + memory_next_offset: None, + memory_refresh_pending: false, } } } @@ -647,9 +703,59 @@ impl OperatorTui { self.selected[Self::screen_index(self.screen)] } + fn memory_request(&self, cursor: Option) -> Option { + let agent = self.snapshot.memory_agents.get(self.memory_agent)?; + if self.filter.query.trim().is_empty() { + Some(ControlRequest::BrowseMemories { + mentra_agent_id: agent.mentra_agent_id.clone(), + cursor, + limit: 50, + filter: self.memory_filter.clone(), + sort: if self.memory_sort == MemorySortView::Relevance { + MemorySortView::Newest + } else { + self.memory_sort + }, + }) + } else { + Some(ControlRequest::SearchMemories { + mentra_agent_id: agent.mentra_agent_id.clone(), + query: self.filter.query.clone(), + page: PageRequest { + cursor: (self.memory_current_offset != 0).then_some(self.memory_current_offset), + limit: 50, + }, + filter: self.memory_filter.clone(), + sort: self.memory_sort, + }) + } + } + + fn current_memory_request(&self) -> Option { + let cursor = self.memory_cursor_stack.last().cloned().flatten(); + self.memory_request(cursor) + } + + fn parse_memory_filter(&mut self) { + let mut query = Vec::new(); + for token in self.filter.query.split_whitespace() { + if let Some(value) = token.strip_prefix("source:") { + self.memory_filter.source = (!value.is_empty()).then(|| value.to_owned()); + } else if let Some(value) = token.strip_prefix("from:") { + self.memory_filter.created_from = value.parse().ok(); + } else if let Some(value) = token.strip_prefix("to:") { + self.memory_filter.created_to = value.parse().ok(); + } else { + query.push(token); + } + } + self.filter.query = query.join(" "); + } + fn visible_rows(&self) -> usize { match self.screen { OperatorScreen::Sessions => self.snapshot.sessions.len(), + OperatorScreen::Memories => self.snapshot.memories.len(), OperatorScreen::Roaming => self.snapshot.schedules.len(), OperatorScreen::Approvals => self.snapshot.approvals.len(), OperatorScreen::Timeline => { @@ -687,6 +793,14 @@ impl OperatorTui { #[must_use] pub fn command_shortcut(&self, key: char) -> Option { + if self.screen == OperatorScreen::Memories && key == 'd' { + let agent = self.snapshot.memory_agents.get(self.memory_agent)?; + let record = self.snapshot.memories.get(self.selected_row())?; + return Some(OperatorCommand::ForgetMemory { + mentra_agent_id: agent.mentra_agent_id.clone(), + record_id: record.record_id.clone(), + }); + } match key { 'p' => Some(OperatorCommand::Pause), 'u' => Some(OperatorCommand::Resume), @@ -737,7 +851,11 @@ impl OperatorTui { .unwrap_or(0); self.screen = OperatorScreen::ALL[(i + 1) % OperatorScreen::ALL.len()]; self.scroll = 0; - TuiAction::None + if self.screen == OperatorScreen::Memories { + TuiAction::Request(ControlRequest::ListMemoryAgents { page: page() }) + } else { + TuiAction::None + } } TuiInput::PreviousScreen => { let i = OperatorScreen::ALL @@ -747,7 +865,11 @@ impl OperatorTui { self.screen = OperatorScreen::ALL [(i + OperatorScreen::ALL.len() - 1) % OperatorScreen::ALL.len()]; self.scroll = 0; - TuiAction::None + if self.screen == OperatorScreen::Memories { + TuiAction::Request(ControlRequest::ListMemoryAgents { page: page() }) + } else { + TuiAction::None + } } TuiInput::ScrollUp => { let index = Self::screen_index(self.screen); @@ -821,7 +943,11 @@ impl OperatorTui { } TuiInput::BeginFilter => { self.filter_editing = true; - self.status = "event/audit filter: type query, Enter applies, Esc clears".into(); + self.status = if self.screen == OperatorScreen::Memories { + "memory search/filter: type query, Enter applies, Esc clears".into() + } else { + "event/audit filter: type query, Enter applies, Esc clears".into() + }; TuiAction::None } TuiInput::FilterCharacter(value) => { @@ -840,14 +966,122 @@ impl OperatorTui { self.filter_editing = false; self.selected[Self::screen_index(self.screen)] = 0; self.status = format!("filter applied: {}", self.filter.query); - TuiAction::None + if self.screen == OperatorScreen::Memories { + self.parse_memory_filter(); + self.memory_cursor_stack.clear(); + self.memory_current_offset = 0; + self.memory_next_offset = None; + self.memory_request(None) + .map_or(TuiAction::None, TuiAction::Request) + } else { + TuiAction::None + } } TuiInput::FilterCancel => { self.filter_editing = false; self.filter.query.clear(); self.status = "filter cleared".into(); + if self.screen == OperatorScreen::Memories { + self.memory_current_offset = 0; + self.memory_sort = match self.memory_sort { + MemorySortView::Relevance => MemorySortView::Newest, + sort => sort, + }; + self.memory_request(None) + .map_or(TuiAction::None, TuiAction::Request) + } else { + TuiAction::None + } + } + TuiInput::MemoryLoad => { + if let (Some(agent), Some(record)) = ( + self.snapshot.memory_agents.get(self.memory_agent), + self.snapshot.memories.get(self.selected_row()), + ) { + TuiAction::Request(ControlRequest::GetMemory { + mentra_agent_id: agent.mentra_agent_id.clone(), + record_id: record.record_id.clone(), + }) + } else { + self.memory_request(None) + .map_or(TuiAction::None, TuiAction::Request) + } + } + TuiInput::MemoryNextAgent => { + if !self.snapshot.memory_agents.is_empty() { + self.memory_agent = (self.memory_agent + 1) % self.snapshot.memory_agents.len(); + } + self.memory_cursor_stack.clear(); + self.memory_current_offset = 0; + self.memory_next_offset = None; + self.memory_request(None) + .map_or(TuiAction::None, TuiAction::Request) + } + TuiInput::MemoryNextPage => { + if self.filter.query.trim().is_empty() { + if let Some(cursor) = self.memory_next_cursor.clone() { + self.memory_cursor_stack.push(Some(cursor.clone())); + return self + .memory_request(Some(cursor)) + .map_or(TuiAction::None, TuiAction::Request); + } + } else if let Some(offset) = self.memory_next_offset { + self.memory_current_offset = offset; + return self + .memory_request(None) + .map_or(TuiAction::None, TuiAction::Request); + } TuiAction::None } + TuiInput::MemoryPreviousPage => { + if self.filter.query.trim().is_empty() { + self.memory_cursor_stack.pop(); + let cursor = self.memory_cursor_stack.last().cloned().flatten(); + self.memory_request(cursor) + .map_or(TuiAction::None, TuiAction::Request) + } else { + self.memory_current_offset = self.memory_current_offset.saturating_sub(50); + self.memory_request(None) + .map_or(TuiAction::None, TuiAction::Request) + } + } + TuiInput::MemoryCycleKind => { + self.memory_filter.kind = match self.memory_filter.kind { + None => Some(MemoryKindView::Episode), + Some(MemoryKindView::Episode) => Some(MemoryKindView::Summary), + Some(MemoryKindView::Summary) => Some(MemoryKindView::Fact), + Some(MemoryKindView::Fact) => None, + }; + self.memory_current_offset = 0; + self.memory_cursor_stack.clear(); + self.memory_request(None) + .map_or(TuiAction::None, TuiAction::Request) + } + TuiInput::MemoryTogglePinned => { + self.memory_filter.pinned = match self.memory_filter.pinned { + None => Some(true), + Some(true) => Some(false), + Some(false) => None, + }; + self.memory_current_offset = 0; + self.memory_cursor_stack.clear(); + self.memory_request(None) + .map_or(TuiAction::None, TuiAction::Request) + } + TuiInput::MemoryCycleSort => { + self.memory_sort = match self.memory_sort { + MemorySortView::Newest => MemorySortView::Oldest, + MemorySortView::Oldest if self.filter.query.trim().is_empty() => { + MemorySortView::Newest + } + MemorySortView::Oldest => MemorySortView::Relevance, + MemorySortView::Relevance => MemorySortView::Newest, + }; + self.memory_current_offset = 0; + self.memory_cursor_stack.clear(); + self.memory_request(None) + .map_or(TuiAction::None, TuiAction::Request) + } TuiInput::Command(command) if command.confirmation() == CommandConfirmation::None => { TuiAction::Execute(command) } @@ -856,10 +1090,13 @@ impl OperatorTui { self.pending = Some(command); TuiAction::None } - TuiInput::Confirm => self - .pending - .take() - .map_or(TuiAction::None, TuiAction::Execute), + TuiInput::Confirm => { + let command = self.pending.take(); + self.memory_refresh_pending = command + .as_ref() + .is_some_and(|command| matches!(command, OperatorCommand::ForgetMemory { .. })); + command.map_or(TuiAction::None, TuiAction::Execute) + } TuiInput::Reject => { self.pending = None; self.status = "command cancelled".into(); @@ -897,6 +1134,7 @@ impl OperatorTui { ControlRequest::Health, ControlRequest::Runtime, ControlRequest::Metrics, + ControlRequest::MemoryHealth, ControlRequest::ListSessions { page: page() }, ControlRequest::ListScheduledJobs { page: page() }, ControlRequest::ListPendingApprovals { page: page() }, @@ -957,8 +1195,13 @@ impl OperatorTui { client: &dyn TuiTransport, command: OperatorCommand, ) -> Result<(), TuiError> { + let refresh_memory = matches!(command, OperatorCommand::ForgetMemory { .. }); self.send_and_apply(client, command.request()).await?; - self.refresh(client).await + self.refresh(client).await?; + if refresh_memory && let Some(request) = self.current_memory_request() { + self.send_and_apply(client, request).await?; + } + Ok(()) } async fn send_and_apply( @@ -990,6 +1233,34 @@ impl OperatorTui { ControlPayload::Health(value) => self.snapshot.health = Some(value), ControlPayload::Metrics(value) => self.snapshot.metrics = Some(value), ControlPayload::Runtime(value) => self.snapshot.runtime = Some(value), + ControlPayload::MemoryHealth(value) => self.snapshot.memory_health = Some(value), + ControlPayload::MemoryAgents(Page { items, .. }) => { + let identity = self + .snapshot + .memory_agents + .get(self.memory_agent) + .map(|agent| agent.mentra_agent_id.clone()); + self.snapshot.memory_agents = items; + self.memory_agent = identity + .and_then(|id| { + self.snapshot + .memory_agents + .iter() + .position(|agent| agent.mentra_agent_id == id) + }) + .unwrap_or(0) + .min(self.snapshot.memory_agents.len().saturating_sub(1)); + } + ControlPayload::Memories(page) => { + self.snapshot.memories = page.items; + self.snapshot.memory_detail = None; + self.memory_next_cursor = page.next_cursor; + self.memory_next_offset = page.next_offset; + self.selected[Self::screen_index(OperatorScreen::Memories)] = 0; + } + ControlPayload::MemoryDetail(detail) => { + self.snapshot.memory_detail = Some(detail); + } ControlPayload::Sessions(Page { items, .. }) => { let index = Self::screen_index(OperatorScreen::Sessions); let identity = self @@ -1205,6 +1476,7 @@ impl OperatorTui { match self.screen { OperatorScreen::Overview => self.draw_overview(frame, chunks[1], color), OperatorScreen::Sessions => self.draw_sessions(frame, chunks[1], color), + OperatorScreen::Memories => self.draw_memories(frame, chunks[1], color), OperatorScreen::QueuesAndBudgets => self.draw_performance(frame, chunks[1], color), OperatorScreen::Roaming => self.draw_roaming(frame, chunks[1], color), OperatorScreen::Approvals => self.draw_approvals(frame, chunks[1], color), @@ -1320,6 +1592,15 @@ impl OperatorTui { ) } OperatorScreen::Sessions => format!("sessions={}", self.snapshot.sessions.len()), + OperatorScreen::Memories => self.snapshot.memory_health.as_ref().map_or_else( + || "memory health unavailable".into(), + |health| { + format!( + "memory agents={} records={}", + health.agent_count, health.record_count + ) + }, + ), OperatorScreen::QueuesAndBudgets => self.snapshot.metrics.as_ref().map_or_else( || "metrics unavailable".into(), |metrics| { @@ -1462,6 +1743,79 @@ impl OperatorTui { ); } + fn draw_memories(&self, frame: &mut Frame<'_>, area: Rect, color: bool) { + let rows = self + .snapshot + .memories + .iter() + .enumerate() + .skip(self.selected_row().saturating_sub(5)) + .map(|(index, item)| { + styled_row( + index == self.selected_row(), + color, + vec![ + item.record_id.clone(), + format!("{:?}", item.kind).to_lowercase(), + item.created_at.to_string(), + if item.pinned { "yes" } else { "no" }.to_owned(), + item.preview.clone(), + ], + ) + }); + let agent = self.snapshot.memory_agents.get(self.memory_agent); + let detail = self.snapshot.memory_detail.as_ref().map_or_else( + || { + vec![ + Line::raw(agent.map_or_else( + || "No operator-visible memory agents".into(), + |agent| format!( + "agent: {} avatar={} records={}", + agent.logical_agent_id, agent.avatar_id, agent.record_count + ), + )), + Line::raw(format!( + "query={:?} kind={:?} pinned={:?} sort={:?}", + self.filter.query, + self.memory_filter.kind, + self.memory_filter.pinned, + self.memory_sort + )), + Line::raw("Enter load/detail / search (source:/from:/to:) g agent n/b page k kind i pinned o sort d forget"), + ] + }, + |detail| { + vec![ + Line::raw(format!("record: {}", detail.record.record_id)), + Line::raw(format!( + "source: {:?} revision={}", + detail.record.source, detail.record.source_revision + )), + Line::raw(format!( + "{}{}", + detail.content, + if detail.content_truncated { " [truncated]" } else { "" } + )), + Line::raw(format!( + "metadata: {}{}", + detail.metadata_json, + if detail.metadata_truncated { " [truncated]" } else { "" } + )), + Line::raw("[d] forget (high-risk confirmation)"), + ] + }, + ); + Self::draw_table_detail( + frame, + area, + color, + "Mentra durable memory", + ["Record", "Kind", "Created", "Pinned", "Preview"], + rows, + detail, + ); + } + fn draw_roaming(&self, frame: &mut Frame<'_>, area: Rect, color: bool) { let rows = self .snapshot @@ -2401,7 +2755,10 @@ pub fn run_preferences_terminal(path: PathBuf, color: bool) -> Result<(), TuiErr }; match app.reduce(input) { TuiAction::Exit => return Ok(()), - TuiAction::None | TuiAction::Refresh | TuiAction::Execute(_) => {} + TuiAction::None + | TuiAction::Refresh + | TuiAction::Execute(_) + | TuiAction::Request(_) => {} } } } @@ -2458,6 +2815,10 @@ pub async fn run_terminal_with_preferences( work_tx.try_send(WorkerRequest::Command(command)).map_err(|_| TuiError("control worker busy".into()))?; app.status = "control command queued".into(); } + TuiAction::Request(request) => { + work_tx.try_send(WorkerRequest::Request(request)).map_err(|_| TuiError("control worker busy".into()))?; + app.status = "memory request queued".into(); + } } } _ = refresh.tick() => { @@ -2471,6 +2832,12 @@ pub async fn run_terminal_with_preferences( app.refresh_millis = result.elapsed_millis; app.sample_performance(); app.status = result.status; + if app.memory_refresh_pending { + app.memory_refresh_pending = false; + if let Some(request) = app.current_memory_request() { + let _ = work_tx.try_send(WorkerRequest::Request(request)); + } + } } Err(error) => { app.record_error(error.to_string()); @@ -2485,6 +2852,7 @@ pub async fn run_terminal_with_preferences( enum WorkerRequest { Refresh, Command(OperatorCommand), + Request(ControlRequest), } struct WorkerResult { @@ -2528,11 +2896,16 @@ async fn worker_cycle( payloads.push(worker_send(client, sequence, command.request()).await?); "command completed".to_owned() } + WorkerRequest::Request(request) => { + payloads.push(worker_send(client, sequence, request).await?); + return Ok((payloads, "memory request completed".to_owned())); + } }; for request in [ ControlRequest::Health, ControlRequest::Runtime, ControlRequest::Metrics, + ControlRequest::MemoryHealth, ControlRequest::ListSessions { page: page() }, ControlRequest::ListScheduledJobs { page: page() }, ControlRequest::ListPendingApprovals { page: page() }, @@ -2611,9 +2984,35 @@ fn map_event(app: &OperatorTui, value: &event::Event) -> Option { event::KeyCode::BackTab | event::KeyCode::Left => Some(TuiInput::PreviousScreen), event::KeyCode::Up => Some(TuiInput::ScrollUp), event::KeyCode::Down => Some(TuiInput::ScrollDown), - event::KeyCode::Char('/') if app.screen == OperatorScreen::Timeline => { + event::KeyCode::Char('/') + if matches!( + app.screen, + OperatorScreen::Timeline | OperatorScreen::Memories + ) => + { Some(TuiInput::BeginFilter) } + event::KeyCode::Enter if app.screen == OperatorScreen::Memories => { + Some(TuiInput::MemoryLoad) + } + event::KeyCode::Char('g') if app.screen == OperatorScreen::Memories => { + Some(TuiInput::MemoryNextAgent) + } + event::KeyCode::Char('n') if app.screen == OperatorScreen::Memories => { + Some(TuiInput::MemoryNextPage) + } + event::KeyCode::Char('b') if app.screen == OperatorScreen::Memories => { + Some(TuiInput::MemoryPreviousPage) + } + event::KeyCode::Char('k') if app.screen == OperatorScreen::Memories => { + Some(TuiInput::MemoryCycleKind) + } + event::KeyCode::Char('i') if app.screen == OperatorScreen::Memories => { + Some(TuiInput::MemoryTogglePinned) + } + event::KeyCode::Char('o') if app.screen == OperatorScreen::Memories => { + Some(TuiInput::MemoryCycleSort) + } event::KeyCode::Char('r') => Some(TuiInput::Refresh), event::KeyCode::Char(key @ ('p' | 'u' | 'f' | 'a' | 'd' | 'x' | 'e' | 't' | 's')) => { app.command_shortcut(key).map(TuiInput::Command) diff --git a/crates/metacrate-grid-agent/src/tui_tests.rs b/crates/metacrate-grid-agent/src/tui_tests.rs index 290fc72..2ca0a15 100644 --- a/crates/metacrate-grid-agent/src/tui_tests.rs +++ b/crates/metacrate-grid-agent/src/tui_tests.rs @@ -80,6 +80,12 @@ impl TuiTransport for FakeTransport { }) } ControlRequest::Metrics => ControlPayload::Metrics(empty_metrics()), + ControlRequest::MemoryHealth => { + ControlPayload::MemoryHealth(crate::MemoryHealthView { + agent_count: 1, + record_count: 2, + }) + } _ => ControlPayload::Completed, }; Ok(ControlResponseEnvelope { @@ -140,7 +146,7 @@ async fn snapshot_refresh_is_transport_neutral_and_bounded() { app.snapshot.health.as_ref().expect("health").service_state, "running" ); - assert_eq!(transport.0.lock().expect("requests").len(), 8); + assert_eq!(transport.0.lock().expect("requests").len(), 9); app.refresh(&transport).await.expect("stable refresh"); assert_eq!( app.snapshot.timeline.len(), @@ -149,6 +155,54 @@ async fn snapshot_refresh_is_transport_neutral_and_bounded() { ); } +#[test] +fn memory_panel_builds_filtered_requests_and_confirms_forget() { + let mut app = OperatorTui::default(); + assert_eq!(app.reduce(TuiInput::NextScreen), TuiAction::None); + assert!(matches!( + app.reduce(TuiInput::NextScreen), + TuiAction::Request(ControlRequest::ListMemoryAgents { .. }) + )); + app.snapshot.memory_agents.push(crate::MemoryAgentView { + avatar_id: "avatar-a".into(), + logical_agent_id: "logical-a".into(), + mentra_agent_id: "mentra-a".into(), + record_count: 1, + }); + app.snapshot.memories.push(crate::MemoryRecordView { + record_id: "fact:1".into(), + kind: crate::MemoryKindView::Fact, + preview: "durable detail".into(), + source: Some("manual".into()), + source_revision: 1, + pinned: true, + created_at: 7, + score: None, + }); + let forget = app.command_shortcut('d').expect("forget command"); + assert_eq!(forget.confirmation(), CommandConfirmation::HighRisk); + + app.reduce(TuiInput::BeginFilter); + for character in "source:manual hello".chars() { + app.reduce(TuiInput::FilterCharacter(character)); + } + let TuiAction::Request(ControlRequest::SearchMemories { query, filter, .. }) = + app.reduce(TuiInput::FilterCommit) + else { + panic!("memory search request") + }; + assert_eq!(query, "hello"); + assert_eq!(filter.source.as_deref(), Some("manual")); + assert!( + app.render(TuiRenderOptions { + width: 120, + height: 30, + color: false + }) + .contains("durable detail") + ); +} + #[tokio::test] async fn performance_history_is_bounded_and_large_dashboards_are_deterministic() { let transport = FakeTransport::new(); @@ -167,6 +221,7 @@ async fn performance_history_is_bounded_and_large_dashboards_are_deterministic() for (screen, marker) in [ (OperatorScreen::Overview, "Runtime"), (OperatorScreen::Sessions, "Sessions"), + (OperatorScreen::Memories, "Mentra durable memory"), (OperatorScreen::QueuesAndBudgets, "Work queue"), (OperatorScreen::Roaming, "Roaming schedules"), (OperatorScreen::Approvals, "Pending approvals"), @@ -239,6 +294,7 @@ fn every_screen_renders_without_a_terminal_at_small_unicode_and_mono_sizes() { for screen in [ OperatorScreen::Overview, OperatorScreen::Sessions, + OperatorScreen::Memories, OperatorScreen::QueuesAndBudgets, OperatorScreen::Roaming, OperatorScreen::Approvals, diff --git a/crates/metacrate-grid-agent/tests/dependency_policy.rs b/crates/metacrate-grid-agent/tests/dependency_policy.rs index ce06f58..7a4e1cc 100644 --- a/crates/metacrate-grid-agent/tests/dependency_policy.rs +++ b/crates/metacrate-grid-agent/tests/dependency_policy.rs @@ -58,10 +58,10 @@ fn package_has_only_reviewed_rust_dependencies_and_no_build_script() { #[test] fn runtime_source_has_no_subprocess_or_native_abi_escape_hatch() { let source = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("src"); - let mut files = Vec::with_capacity(38); + let mut files = Vec::with_capacity(39); collect_rust_files(&source, &mut files); assert!( - files.len() <= 38, + files.len() <= 39, "source-file count needs a reviewed bound update" ); for path in files { diff --git a/docs/grid-agent-control-plane.md b/docs/grid-agent-control-plane.md index e5cf865..36891f1 100644 --- a/docs/grid-agent-control-plane.md +++ b/docs/grid-agent-control-plane.md @@ -48,15 +48,27 @@ The JSON request envelope is stable and versioned. For example: {"version":1,"request_id":"health-1","request":{"method":"health"}} ``` -Observers can call `health`, `metrics`, `runtime`, `list_sessions`, +Observers can call `health`, `metrics`, `runtime`, `memory_health`, `list_sessions`, `list_scheduled_jobs`, `list_pending_approvals`, `list_audit_events`, `list_observability_events`, and `subscribe_events`. Operators can additionally call `cancel_request`, `pause_autonomy`, `resume_autonomy`, `cancel_action`, `decide_approval`, `force_reconnect`, `expire_conversation`, `set_roaming_job`, -`inject_operator_message`, and `graceful_shutdown`. Cancellation is a mutation +`list_memory_agents`, `browse_memories`, `search_memories`, `get_memory`, +`forget_memory`, `inject_operator_message`, and `graceful_shutdown`. Memory +identity and content operations enforce operator role in the runtime target as +well as at the mutation gate; observers receive aggregate memory health only. +Cancellation is a mutation and is operator-only. List requests use an opaque numeric cursor and a page size of 1 through 100. +Memory agents come from the persisted authenticated-avatar mapping written +when an authorized IM agent is used; names are never parsed to recover that +association. Browsing uses Mentra-owned stable `(created_at, record_id)` +keyset cursors, filters and tombstones. Queries, pages, previews, detail content, +metadata, concurrent requests, and response frames remain bounded. Forget uses +Mentra's agent-scoped tombstone operation and audit records contain identifiers +and outcomes, never memory content. + Runtime projections contain lifecycle/readiness, session generation, safe region and pose fields when known, behavior mode, control-queue utilization, and aggregate budget use. Conversation responses contain metadata only. Audit diff --git a/docs/grid-agent-llm.md b/docs/grid-agent-llm.md index 8b6ee5d..042016f 100644 --- a/docs/grid-agent-llm.md +++ b/docs/grid-agent-llm.md @@ -36,6 +36,9 @@ learning without displacing the command being handled. Runtime records use `HybridRuntimeStore`: conversation/runtime state is stored in `runtime.sqlite`, with the associated Mentra memory store and transcript, task, team, and workspace paths under the configured state directory. +`memory-agents.json` atomically persists the authenticated avatar, logical +agent key, and actual Mentra agent ID for operator browsing. This explicit +mapping is restored independently of generated agent names. ## Tools and autonomous safety review diff --git a/docs/grid-agent-tui.md b/docs/grid-agent-tui.md index 0a02125..4635ee4 100644 --- a/docs/grid-agent-tui.md +++ b/docs/grid-agent-tui.md @@ -36,10 +36,18 @@ The Overview dashboard shows readiness, region/pose, behavior, build and visual work, approvals, and recent failures. Queues & budgets shows real queue and rate-limit gauges, active work, outcome counts, dropped events, aggregate latency, and a bounded 120-refresh history. Sessions, roaming schedules, -approvals, timeline/audit, errors, diagnostics, health, and preferences retain +approvals, durable memory, timeline/audit, errors, diagnostics, health, and preferences retain focused tables or panels and contextual detail. Missing backend data is marked unavailable rather than estimated. +The Memory panel is operator-only. Enter loads the selected agent or record, +`g` changes the explicitly mapped avatar agent, `n`/`b` page, `k` cycles record +kind, `i` cycles pinned state, `o` changes sort, and `d` confirms a durable +tombstone. `/` searches content; `source:value`, `from:unix-seconds`, and +`to:unix-seconds` tokens set provenance and creation-time filters. Empty search +text browses all non-tombstoned records. Observer clients receive aggregate +agent/record health only and never memory identity, preview, or content. + Use Tab/Shift-Tab or Left/Right to change panels and Up/Down to select rows. `/` searches event, audit, correlation, time, duration, result, and reason data; Enter applies the search and Esc clears it. `r` refreshes; `p`/`u` pause or @@ -60,9 +68,11 @@ when the bounded worker queue is full. The terminal guard restores raw mode, cursor visibility, and the alternate screen on success, error, Ctrl-C, or panic unwinding. -Diagnostic panels show only explicitly captured redacted envelopes. Prompt -and response content, API keys, grid passwords, operator tokens, capability -URLs, and model reasoning have no TUI representation. +Diagnostic panels show only explicitly captured redacted envelopes. Memory +content appears only in the operator-only Memory detail panel; it is never +copied into diagnostics, audit events, or observer responses. API keys, grid +passwords, operator tokens, capability URLs, and model reasoning have no TUI +representation. Focused verification: diff --git a/vendor/mentra/Cargo.lock b/vendor/mentra/Cargo.lock new file mode 100644 index 0000000..c90cf32 --- /dev/null +++ b/vendor/mentra/Cargo.lock @@ -0,0 +1,2082 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "aho-corasick" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddd31a130427c27518df266943a5308ed92d4b226cc639f5a8f1002816174301" +dependencies = [ + "memchr", +] + +[[package]] +name = "anes" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4b46cbb362ab8752921c97e041f5e366ee6297bd428a31275b9fcf1e380f7299" + +[[package]] +name = "anstyle" +version = "1.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" + +[[package]] +name = "async-trait" +version = "0.1.89" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "atomic-waker" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" + +[[package]] +name = "autocfg" +version = "1.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" + +[[package]] +name = "base64" +version = "0.22.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" + +[[package]] +name = "bitflags" +version = "2.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b588b76d00fde79687d7646a9b5bdf3cc0f655e0bbd080335a95d7e96f3587da" + +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + +[[package]] +name = "bumpalo" +version = "3.20.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" + +[[package]] +name = "bytes" +version = "1.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" + +[[package]] +name = "cast" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5" + +[[package]] +name = "cc" +version = "1.2.67" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e17dd265a7d0f31ef544e1b20e03add05d3b45b491b633b10d67145d2acc1a38" +dependencies = [ + "find-msvc-tools", + "jobserver", + "libc", + "shlex", +] + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "cfg_aliases" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" + +[[package]] +name = "chacha20" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "rand_core 0.10.1", +] + +[[package]] +name = "ciborium" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42e69ffd6f0917f5c029256a24d0161db17cea3997d185db0d35926308770f0e" +dependencies = [ + "ciborium-io", + "ciborium-ll", + "serde", +] + +[[package]] +name = "ciborium-io" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "05afea1e0a06c9be33d539b876f1ce3692f4afea2cb41f740e7743225ed1c757" + +[[package]] +name = "ciborium-ll" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57663b653d948a338bfb3eeba9bb2fd5fcfaecb9e199e87e1eda4d9e8b240fd9" +dependencies = [ + "ciborium-io", + "half", +] + +[[package]] +name = "clap" +version = "4.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dd059f9da4f5c36b3787f65d38ccaab1cc315f07b01f89abc8359ee6a8205011" +dependencies = [ + "clap_builder", +] + +[[package]] +name = "clap_builder" +version = "4.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f09628afdcc538b57f3c6341e9c8e9970f18e4a481690a64974d7023bd33548b" +dependencies = [ + "anstyle", + "clap_lex", +] + +[[package]] +name = "clap_lex" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" + +[[package]] +name = "cpufeatures" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +dependencies = [ + "libc", +] + +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + +[[package]] +name = "criterion" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2b12d017a929603d80db1831cd3a24082f8137ce19c69e6447f54f5fc8d692f" +dependencies = [ + "anes", + "cast", + "ciborium", + "clap", + "criterion-plot", + "futures", + "is-terminal", + "itertools", + "num-traits", + "once_cell", + "oorandom", + "plotters", + "rayon", + "regex", + "serde", + "serde_derive", + "serde_json", + "tinytemplate", + "tokio", + "walkdir", +] + +[[package]] +name = "criterion-plot" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b50826342786a51a89e2da3a28f1c32b06e387201bc2d19791f622c673706b1" +dependencies = [ + "cast", + "itertools", +] + +[[package]] +name = "crossbeam-deque" +version = "0.8.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5181e0de7b61eb03a81e347d6dd8797bae9da5146707b51077e2d71a54ec0ceb" +dependencies = [ + "crossbeam-epoch", + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-epoch" +version = "0.9.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d6914041f254d6e9176c01941b21115dcfb7089e55135a35411081bd106ef3f" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17" + +[[package]] +name = "crunchy" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" + +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "typenum", +] + +[[package]] +name = "data-encoding" +version = "2.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4ae5f15dda3c708c0ade84bfee31ccab44a3da4f88015ed22f63732abe300c8" + +[[package]] +name = "deranged" +version = "0.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c" +dependencies = [ + "powerfmt", + "serde_core", +] + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", +] + +[[package]] +name = "directories" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "16f5094c54661b38d03bd7e50df373292118db60b585c08a411c6d840017fe7d" +dependencies = [ + "dirs-sys", +] + +[[package]] +name = "dirs-sys" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e01a3366d27ee9890022452ee61b2b63a67e6f13f58900b651ff5665f0bb1fab" +dependencies = [ + "libc", + "option-ext", + "redox_users", + "windows-sys 0.61.2", +] + +[[package]] +name = "either" +version = "1.16.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "fallible-iterator" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2acce4a10f12dc2fb14a218589d4f1f62ef011b2d0cc4b3cb1bba8e94da14649" + +[[package]] +name = "fallible-streaming-iterator" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7360491ce676a36bf9bb3c56c1aa791658183a54d2744120f27285738d90465a" + +[[package]] +name = "find-msvc-tools" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" + +[[package]] +name = "foldhash" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb" + +[[package]] +name = "form_urlencoded" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" +dependencies = [ + "percent-encoding", +] + +[[package]] +name = "futures" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d" +dependencies = [ + "futures-channel", + "futures-core", + "futures-io", + "futures-sink", + "futures-task", + "futures-util", +] + +[[package]] +name = "futures-channel" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" +dependencies = [ + "futures-core", + "futures-sink", +] + +[[package]] +name = "futures-core" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" + +[[package]] +name = "futures-io" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" + +[[package]] +name = "futures-macro" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "futures-sink" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893" + +[[package]] +name = "futures-task" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393" + +[[package]] +name = "futures-util" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" +dependencies = [ + "futures-core", + "futures-io", + "futures-macro", + "futures-sink", + "futures-task", + "memchr", + "pin-project-lite", + "slab", +] + +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + +[[package]] +name = "getrandom" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" +dependencies = [ + "cfg-if", + "js-sys", + "libc", + "wasi", + "wasm-bindgen", +] + +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi 5.3.0", + "wasip2", +] + +[[package]] +name = "getrandom" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "300e883d756b2e4ec94e02791f39b04b522276138852cfc41d9fb7e904106099" +dependencies = [ + "cfg-if", + "js-sys", + "libc", + "r-efi 6.0.0", + "rand_core 0.10.1", + "wasm-bindgen", +] + +[[package]] +name = "glob-match" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9985c9503b412198aa4197559e9a318524ebc4519c229bfa05a535828c950b9d" + +[[package]] +name = "half" +version = "2.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" +dependencies = [ + "cfg-if", + "crunchy", + "zerocopy", +] + +[[package]] +name = "hashbrown" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" +dependencies = [ + "foldhash", +] + +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" + +[[package]] +name = "hashlink" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "824e001ac4f3012dd16a264bec811403a67ca9deb6c102fc5049b32c4574b35f" +dependencies = [ + "hashbrown 0.16.1", +] + +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + +[[package]] +name = "hermit-abi" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" + +[[package]] +name = "http" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6970f50e31d6fc17d3fa27329444bfa74e196cf62e95052a3f6fee181dba6425" +dependencies = [ + "bytes", + "itoa", +] + +[[package]] +name = "http-body" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ca2a8f2913ee65f60facd6a5905613afaa448497a0230cc41ce022d93290bc2c" +dependencies = [ + "bytes", + "http", +] + +[[package]] +name = "http-body-util" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e9f41fd6a08e4d4ec69df65976da761afd5ad5e58a9d4acb46bd1c953a9e3ff2" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "pin-project-lite", +] + +[[package]] +name = "httparse" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" + +[[package]] +name = "hyper" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "55281c53a1894c864990125767da440a4e630446785086f52523b20033b74498" +dependencies = [ + "atomic-waker", + "bytes", + "futures-channel", + "futures-core", + "http", + "http-body", + "httparse", + "itoa", + "pin-project-lite", + "smallvec", + "tokio", + "want", +] + +[[package]] +name = "hyper-rustls" +version = "0.27.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "33ca68d021ef39cf6463ab54c1d0f5daf03377b70561305bb89a8f83aab66e0f" +dependencies = [ + "http", + "hyper", + "hyper-util", + "rustls", + "tokio", + "tokio-rustls", + "tower-service", + "webpki-roots", +] + +[[package]] +name = "hyper-util" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" +dependencies = [ + "base64", + "bytes", + "futures-channel", + "futures-util", + "http", + "http-body", + "hyper", + "ipnet", + "libc", + "percent-encoding", + "pin-project-lite", + "socket2", + "tokio", + "tower-service", + "tracing", +] + +[[package]] +name = "idna" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "634d9b1461af396cad843f47fdba5597a4f9e6ddd4bfb6ff5d85028c25cb12f6" +dependencies = [ + "unicode-bidi", + "unicode-normalization", +] + +[[package]] +name = "indexmap" +version = "2.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" +dependencies = [ + "equivalent", + "hashbrown 0.17.1", +] + +[[package]] +name = "ipnet" +version = "2.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" + +[[package]] +name = "is-terminal" +version = "0.4.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" +dependencies = [ + "hermit-abi", + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "itertools" +version = "0.10.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b0fd2260e829bddf4cb6ea802289de2f86d6a7a690192fbe91b3f46e0f2c8473" +dependencies = [ + "either", +] + +[[package]] +name = "itoa" +version = "1.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" + +[[package]] +name = "jobserver" +version = "0.1.35" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1c00acbd29eabad4a2392fa0e921c874934dbbf4194312ad20f04a0ed67a3cb3" +dependencies = [ + "getrandom 0.4.3", + "libc", +] + +[[package]] +name = "js-sys" +version = "0.3.99" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "142bc4740e452c1e57ade0cbc129f139c9093e354346f0872ef985f4f5cf5f11" +dependencies = [ + "cfg-if", + "futures-util", + "once_cell", + "wasm-bindgen", +] + +[[package]] +name = "libc" +version = "0.2.186" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" + +[[package]] +name = "libredox" +version = "0.1.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f02ab6bace2054fb888a3c16f990117b579d14a3088e472d63c6011fa185c9d3" +dependencies = [ + "libc", +] + +[[package]] +name = "libsqlite3-sys" +version = "0.37.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b1f111c8c41e7c61a49cd34e44c7619462967221a6443b0ec299e0ac30cfb9b1" +dependencies = [ + "cc", + "pkg-config", + "vcpkg", +] + +[[package]] +name = "lock_api" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" +dependencies = [ + "scopeguard", +] + +[[package]] +name = "log" +version = "0.4.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" + +[[package]] +name = "lru-slab" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" + +[[package]] +name = "memchr" +version = "2.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" + +[[package]] +name = "mentra" +version = "0.18.3" +dependencies = [ + "async-trait", + "base64", + "criterion", + "directories", + "futures-util", + "glob-match", + "libc", + "mentra-provider", + "rand 0.9.5", + "regex", + "reqwest", + "ring", + "rusqlite", + "serde", + "serde_json", + "serde_yaml_ng", + "similar", + "strum", + "thiserror", + "time", + "tokio", + "unicode-normalization", + "url", + "windows-sys 0.61.2", +] + +[[package]] +name = "mentra-provider" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e9f644037cf45558b393becc1ac080f0d55edced3093a155a7cd3431c1180b7f" +dependencies = [ + "async-trait", + "base64", + "futures-util", + "http", + "reqwest", + "serde", + "serde_json", + "strum", + "thiserror", + "time", + "tokio", + "tokio-tungstenite", + "url", + "zstd", +] + +[[package]] +name = "mio" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "30d65c71f1ce40ab09135ce117d742b9f8a19ff91a41a8b57ed50bc2de59c427" +dependencies = [ + "libc", + "wasi", + "windows-sys 0.61.2", +] + +[[package]] +name = "num-conv" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "51d515d32fb182ee37cda2ccdcb92950d6a3c2893aa280e540671c2cd0f3b1d9" + +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", +] + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "oorandom" +version = "11.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e" + +[[package]] +name = "option-ext" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d" + +[[package]] +name = "parking_lot" +version = "0.12.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93857453250e3077bd71ff98b6a65ea6621a19bb0f559a85248955ac12c45a1a" +dependencies = [ + "lock_api", + "parking_lot_core", +] + +[[package]] +name = "parking_lot_core" +version = "0.9.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1" +dependencies = [ + "cfg-if", + "libc", + "redox_syscall", + "smallvec", + "windows-link", +] + +[[package]] +name = "percent-encoding" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" + +[[package]] +name = "pin-project-lite" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" + +[[package]] +name = "pkg-config" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" + +[[package]] +name = "plotters" +version = "0.3.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5aeb6f403d7a4911efb1e33402027fc44f29b5bf6def3effcc22d7bb75f2b747" +dependencies = [ + "num-traits", + "plotters-backend", + "plotters-svg", + "wasm-bindgen", + "web-sys", +] + +[[package]] +name = "plotters-backend" +version = "0.3.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df42e13c12958a16b3f7f4386b9ab1f3e7933914ecea48da7139435263a4172a" + +[[package]] +name = "plotters-svg" +version = "0.3.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "51bae2ac328883f7acdfea3d66a7c35751187f870bc81f94563733a154d7a670" +dependencies = [ + "plotters-backend", +] + +[[package]] +name = "powerfmt" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "439ee305def115ba05938db6eb1644ff94165c5ab5e9420d1c1bcedbba909391" + +[[package]] +name = "ppv-lite86" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" +dependencies = [ + "zerocopy", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quinn" +version = "0.11.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c1a41e437b6bbd489372cd4971de128e85c855f56c57f283d20ff016cf7c0a8" +dependencies = [ + "bytes", + "cfg_aliases", + "pin-project-lite", + "quinn-proto", + "quinn-udp", + "rustc-hash", + "rustls", + "socket2", + "thiserror", + "tokio", + "tracing", + "web-time", +] + +[[package]] +name = "quinn-proto" +version = "0.11.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2f4bfc015262b9df63c8845072ce59068853ff5872180c2ce2f13038b970e560" +dependencies = [ + "bytes", + "getrandom 0.4.3", + "lru-slab", + "rand 0.10.2", + "rand_pcg", + "ring", + "rustc-hash", + "rustls", + "rustls-pki-types", + "slab", + "thiserror", + "tinyvec", + "tracing", + "web-time", +] + +[[package]] +name = "quinn-udp" +version = "0.5.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "35a133f956daabe89a61a685c2649f13d82d5aa4bd5d12d1277e1072a21c0694" +dependencies = [ + "cfg_aliases", + "libc", + "once_cell", + "socket2", + "tracing", + "windows-sys 0.61.2", +] + +[[package]] +name = "quote" +version = "1.0.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + +[[package]] +name = "r-efi" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" + +[[package]] +name = "rand" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" +dependencies = [ + "rand_chacha", + "rand_core 0.9.5", +] + +[[package]] +name = "rand" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" +dependencies = [ + "chacha20", + "getrandom 0.4.3", + "rand_core 0.10.1", +] + +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core 0.9.5", +] + +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", +] + +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + +[[package]] +name = "rand_pcg" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "caa0f4137e1c0a72f4c651489402276c8e8e1cf081f3b0ba156d2cbeef09e86a" +dependencies = [ + "rand_core 0.10.1", +] + +[[package]] +name = "rayon" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fb39b166781f92d482534ef4b4b1b2568f42613b53e5b6c160e24cfbfa30926d" +dependencies = [ + "either", + "rayon-core", +] + +[[package]] +name = "rayon-core" +version = "1.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22e18b0f0062d30d4230b2e85ff77fdfe4326feb054b9783a3460d8435c8ab91" +dependencies = [ + "crossbeam-deque", + "crossbeam-utils", +] + +[[package]] +name = "redox_syscall" +version = "0.5.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" +dependencies = [ + "bitflags", +] + +[[package]] +name = "redox_users" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4e608c6638b9c18977b00b475ac1f28d14e84b27d8d42f70e0bf1e3dec127ac" +dependencies = [ + "getrandom 0.2.17", + "libredox", + "thiserror", +] + +[[package]] +name = "regex" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f020237b6c8eed93db2e2cb53c00c60a8e1bc73da7d073199a1180401450218d" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fcfdb36bda0c880c5931cdc7a2bcdc8ba4556847b9d912bca70bc94708711ad" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + +[[package]] +name = "reqwest" +version = "0.12.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" +dependencies = [ + "base64", + "bytes", + "futures-core", + "futures-util", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-rustls", + "hyper-util", + "js-sys", + "log", + "percent-encoding", + "pin-project-lite", + "quinn", + "rustls", + "rustls-pki-types", + "serde", + "serde_json", + "serde_urlencoded", + "sync_wrapper", + "tokio", + "tokio-rustls", + "tokio-util", + "tower", + "tower-http", + "tower-service", + "url", + "wasm-bindgen", + "wasm-bindgen-futures", + "wasm-streams", + "web-sys", + "webpki-roots", +] + +[[package]] +name = "ring" +version = "0.17.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7" +dependencies = [ + "cc", + "cfg-if", + "getrandom 0.2.17", + "libc", + "untrusted", + "windows-sys 0.52.0", +] + +[[package]] +name = "rsqlite-vfs" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c51c9ae4df8a7fba42103df5c621fa3c37eccf3a3c650879e90fc48b11cc192c" +dependencies = [ + "hashbrown 0.16.1", + "thiserror", +] + +[[package]] +name = "rusqlite" +version = "0.39.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a0d2b0146dd9661bf67bb107c0bb2a55064d556eeb3fc314151b957f313bcd4e" +dependencies = [ + "bitflags", + "fallible-iterator", + "fallible-streaming-iterator", + "hashlink", + "libsqlite3-sys", + "smallvec", + "sqlite-wasm-rs", +] + +[[package]] +name = "rustc-hash" +version = "2.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" + +[[package]] +name = "rustls" +version = "0.23.42" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3c54fcab019b409d04215d3a17cb438fd7fbf192ee61461f20f4fe18704bc138" +dependencies = [ + "once_cell", + "ring", + "rustls-pki-types", + "rustls-webpki", + "subtle", + "zeroize", +] + +[[package]] +name = "rustls-pki-types" +version = "1.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "764899a24af3980067ee14bc143654f297b22eaebfe3c7b6b211920a5a59b046" +dependencies = [ + "web-time", + "zeroize", +] + +[[package]] +name = "rustls-webpki" +version = "0.103.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", +] + +[[package]] +name = "rustversion" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" + +[[package]] +name = "ryu" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" + +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + +[[package]] +name = "scopeguard" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.150" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] +name = "serde_urlencoded" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3491c14715ca2294c4d6a88f15e84739788c1d030eed8c110436aafdaa2f3fd" +dependencies = [ + "form_urlencoded", + "itoa", + "ryu", + "serde", +] + +[[package]] +name = "serde_yaml_ng" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b4db627b98b36d4203a7b458cf3573730f2bb591b28871d916dfa9efabfd41f" +dependencies = [ + "indexmap", + "itoa", + "ryu", + "serde", + "unsafe-libyaml", +] + +[[package]] +name = "sha1" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest", +] + +[[package]] +name = "shlex" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba" + +[[package]] +name = "signal-hook-registry" +version = "1.4.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4db69cba1110affc0e9f7bcd48bbf87b3f4fc7c61fc9155afd4c469eb3d6c1b" +dependencies = [ + "errno", + "libc", +] + +[[package]] +name = "similar" +version = "2.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbbb5d9659141646ae647b42fe094daf6c6192d1620870b449d9557f748b2daa" + +[[package]] +name = "slab" +version = "0.4.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c790de23124f9ab44544d7ac05d60440adc586479ce501c1d6d7da3cd8c9cf5" + +[[package]] +name = "smallvec" +version = "1.15.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" + +[[package]] +name = "socket2" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3d1e2c7f27f8d4cb10542a02c49005dbd6e93095799d6f3be745fae9f8fedd4" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "sqlite-wasm-rs" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc3efc0da82635d7e1ced0053bbbfa8c7ab9645d0bf36ceb4f7127bb85315d75" +dependencies = [ + "cc", + "js-sys", + "rsqlite-vfs", + "wasm-bindgen", +] + +[[package]] +name = "strum" +version = "0.27.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af23d6f6c1a224baef9d3f61e287d2761385a5b88fdab4eb4c6f11aeb54c4bcf" +dependencies = [ + "strum_macros", +] + +[[package]] +name = "strum_macros" +version = "0.27.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7695ce3845ea4b33927c055a39dc438a45b059f7c1b3d91d38d10355fb8cbca7" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "subtle" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" + +[[package]] +name = "syn" +version = "2.0.119" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "sync_wrapper" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" +dependencies = [ + "futures-core", +] + +[[package]] +name = "thiserror" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "time" +version = "0.3.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f9e442fc33d7fdb45aa9bfeb312c095964abdf596f7567261062b2a7107aaabd" +dependencies = [ + "deranged", + "itoa", + "num-conv", + "powerfmt", + "serde_core", + "time-core", + "time-macros", +] + +[[package]] +name = "time-core" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b36ee98fd31ec7426d599183e8fe26932a8dc1fb76ddb6214d05493377d34ca" + +[[package]] +name = "time-macros" +version = "0.2.25" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "71e552d1249bf61ac2a52db88179fd0673def1e1ad8243a00d9ec9ed71fee3dd" +dependencies = [ + "num-conv", + "time-core", +] + +[[package]] +name = "tinytemplate" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be4d6b5f19ff7664e8c98d03e2139cb510db9b0a60b55f8e8709b689d939b6bc" +dependencies = [ + "serde", + "serde_json", +] + +[[package]] +name = "tinyvec" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb4ebadaa0af04fab11ae01eb5f9fdb5f9c5b875506e210e71c07873528baa7f" +dependencies = [ + "tinyvec_macros", +] + +[[package]] +name = "tinyvec_macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" + +[[package]] +name = "tokio" +version = "1.52.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" +dependencies = [ + "bytes", + "libc", + "mio", + "parking_lot", + "pin-project-lite", + "signal-hook-registry", + "socket2", + "tokio-macros", + "windows-sys 0.61.2", +] + +[[package]] +name = "tokio-macros" +version = "2.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tokio-rustls" +version = "0.26.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" +dependencies = [ + "rustls", + "tokio", +] + +[[package]] +name = "tokio-tungstenite" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d25a406cddcc431a75d3d9afc6a7c0f7428d4891dd973e4d54c56b46127bf857" +dependencies = [ + "futures-util", + "log", + "tokio", + "tungstenite", +] + +[[package]] +name = "tokio-util" +version = "0.7.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ae9cec805b01e8fc3fd2fe289f89149a9b66dd16786abd8b19cfa7b48cb0098" +dependencies = [ + "bytes", + "futures-core", + "futures-sink", + "pin-project-lite", + "tokio", +] + +[[package]] +name = "tower" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" +dependencies = [ + "futures-core", + "futures-util", + "pin-project-lite", + "sync_wrapper", + "tokio", + "tower-layer", + "tower-service", +] + +[[package]] +name = "tower-http" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" +dependencies = [ + "bitflags", + "bytes", + "futures-util", + "http", + "http-body", + "pin-project-lite", + "tower", + "tower-layer", + "tower-service", + "url", +] + +[[package]] +name = "tower-layer" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "121c2a6cda46980bb0fcd1647ffaf6cd3fc79a013de288782836f6df9c48780e" + +[[package]] +name = "tower-service" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8df9b6e13f2d32c91b9bd719c00d1958837bc7dec474d94952798cc8e69eeec3" + +[[package]] +name = "tracing" +version = "0.1.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" +dependencies = [ + "pin-project-lite", + "tracing-core", +] + +[[package]] +name = "tracing-core" +version = "0.1.36" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" +dependencies = [ + "once_cell", +] + +[[package]] +name = "try-lock" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" + +[[package]] +name = "tungstenite" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8628dcc84e5a09eb3d8423d6cb682965dea9133204e8fb3efee74c2a0c259442" +dependencies = [ + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "rand 0.9.5", + "sha1", + "thiserror", + "utf-8", +] + +[[package]] +name = "typenum" +version = "1.20.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" + +[[package]] +name = "unicode-bidi" +version = "0.3.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c1cb5db39152898a79168971543b1cb5020dff7fe43c8dc468b0885f5e29df5" + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "unicode-normalization" +version = "0.1.25" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5fd4f6878c9cb28d874b009da9e8d183b5abc80117c40bbd187a1fde336be6e8" +dependencies = [ + "tinyvec", +] + +[[package]] +name = "unsafe-libyaml" +version = "0.2.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "673aac59facbab8a9007c7f6108d11f63b603f7cabff99fabf650fea5c32b861" + +[[package]] +name = "untrusted" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" + +[[package]] +name = "url" +version = "2.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22784dbdf76fdde8af1aeda5622b546b422b6fc585325248a2bf9f5e41e94d6c" +dependencies = [ + "form_urlencoded", + "idna", + "percent-encoding", +] + +[[package]] +name = "utf-8" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" + +[[package]] +name = "vcpkg" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" + +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + +[[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + +[[package]] +name = "want" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bfa7760aed19e106de2c7c0b581b509f2f25d3dacaf737cb82ac61bc6d760b0e" +dependencies = [ + "try-lock", +] + +[[package]] +name = "wasi" +version = "0.11.1+wasi-snapshot-preview1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" + +[[package]] +name = "wasip2" +version = "1.0.3+wasi-0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "20064672db26d7cdc89c7798c48a0fdfac8213434a1186e5ef29fd560ae223d6" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasm-bindgen" +version = "0.2.122" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3ed04576f974d2b2fba0f38c51dbc5518011e38c36bf1143164be765528fd409" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-futures" +version = "0.4.72" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9473dbd2991ae90b6291c3c32c30c6187ac49aa32f9905d1cce280ec1e110b0f" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.122" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "916151b09da36bd82f6615cbf3a419e2f0ba23a03c6160e8e92eb6bd4aa1dec6" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.122" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "299047362ccbfce148b67ab7e73349f77748e00c8296f9542adfad2ad82c5c5e" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.122" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a929b2c61f11ba3e9bc35b50c1f25cb38e0e892c0c231ae2b8cf78d5dad4437" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "wasm-streams" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "15053d8d85c7eccdbefef60f06769760a563c7f0a9d6902a13d35c7800b0ad65" +dependencies = [ + "futures-util", + "js-sys", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "web-sys" +version = "0.3.99" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6d621441cfc37b84979402712047321980c178f299193a3589d05b99e8763436" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "web-time" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a6580f308b1fad9207618087a65c04e7a10bc77e02c8e84e9b00dd4b12fa0bb" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "webpki-roots" +version = "1.0.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bf85cb06032201fa7c6f829d7db5a7e5aa45bcc0655327713065f6f0576731bf" +dependencies = [ + "rustls-pki-types", +] + +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets", +] + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + +[[package]] +name = "wit-bindgen" +version = "0.57.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" + +[[package]] +name = "zerocopy" +version = "0.8.54" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7cbbc0a705a0fd05cc3676525980d2bf5a9bc4adac6d6475209a7887cf59d19" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.54" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e2e817b7b52d0c7358d3246da9d69935ebb18116b2b102b4230dac079b4862f5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "zeroize" +version = "1.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" + +[[package]] +name = "zmij" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" + +[[package]] +name = "zstd" +version = "0.13.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e91ee311a569c327171651566e07972200e76fcfe2242a4fa446149a3881c08a" +dependencies = [ + "zstd-safe", +] + +[[package]] +name = "zstd-safe" +version = "7.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f49c4d5f0abb602a93fb8736af2a4f4dd9512e36f7f570d66e65ff867ed3b9d" +dependencies = [ + "zstd-sys", +] + +[[package]] +name = "zstd-sys" +version = "2.0.16+zstd.1.5.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91e19ebc2adc8f83e43039e79776e3fda8ca919132d68a1fed6a5faca2683748" +dependencies = [ + "cc", + "pkg-config", +] diff --git a/vendor/mentra/Cargo.toml b/vendor/mentra/Cargo.toml new file mode 100644 index 0000000..109244f --- /dev/null +++ b/vendor/mentra/Cargo.toml @@ -0,0 +1,177 @@ +# THIS FILE IS AUTOMATICALLY GENERATED BY CARGO +# +# When uploading crates to the registry Cargo will automatically +# "normalize" Cargo.toml files for maximal compatibility +# with all versions of Cargo and also rewrite `path` dependencies +# to registry (e.g., crates.io) dependencies. +# +# If you are reading this file be aware that the original Cargo.toml +# will likely look very different (and much more reasonable). +# See Cargo.toml.orig for the original contents. + +[package] +edition = "2024" +rust-version = "1.88" +name = "mentra" +version = "0.18.3" +build = false +autolib = false +autobins = false +autoexamples = false +autotests = false +autobenches = false +description = "An agent runtime for tool-using LLM applications" +homepage = "https://github.com/oops-rs/mentra" +documentation = "https://docs.rs/mentra" +readme = "README.md" +keywords = [ + "agents", + "llm", + "runtime", + "tools", + "ai", +] +categories = [ + "asynchronous", + "development-tools", +] +license = "MIT" +repository = "https://github.com/oops-rs/mentra" + +[features] +default = ["responses-websocket"] +openai-oauth = ["dep:ring"] +responses-websocket = ["mentra-provider/responses-websocket"] +test-utils = [] + +[lib] +name = "mentra" +path = "src/lib.rs" + +[[test]] +name = "agent_runtime" +path = "tests/agent_runtime.rs" + +[[test]] +name = "branching" +path = "tests/branching.rs" + +[[test]] +name = "mcp_sse_smoke" +path = "tests/mcp_sse_smoke.rs" + +[[test]] +name = "public_api" +path = "tests/public_api.rs" + +[[test]] +name = "responses_transport" +path = "tests/responses_transport.rs" + +[[test]] +name = "skills_api" +path = "tests/skills_api.rs" + +[[bench]] +name = "long_session" +path = "benches/long_session.rs" +harness = false +required-features = ["test-utils"] + +[dependencies.async-trait] +version = "0.1.89" + +[dependencies.base64] +version = "0.22.1" + +[dependencies.directories] +version = "6.0.0" + +[dependencies.futures-util] +version = "0.3.31" + +[dependencies.glob-match] +version = "0.2" + +[dependencies.libc] +version = "0.2" + +[dependencies.mentra-provider] +version = "0.5.1" +default-features = false + +[dependencies.rand] +version = "0.9.2" + +[dependencies.regex] +version = "1.12.2" + +[dependencies.reqwest] +version = "0.12.23" +features = [ + "json", + "rustls-tls", + "stream", +] +default-features = false + +[dependencies.ring] +version = "0.17.14" +optional = true + +[dependencies.rusqlite] +version = "0.39" +features = ["bundled"] + +[dependencies.serde] +version = "1.0.228" +features = ["derive"] + +[dependencies.serde_json] +version = "1.0.149" + +[dependencies.serde_yaml_ng] +version = "0.10" + +[dependencies.similar] +version = "2.7.0" + +[dependencies.strum] +version = "0.27" +features = ["derive"] + +[dependencies.thiserror] +version = "2.0.18" + +[dependencies.time] +version = "0.3" +features = [ + "formatting", + "parsing", + "serde", +] + +[dependencies.tokio] +version = "1.50.0" +features = ["full"] + +[dependencies.unicode-normalization] +version = "0.1.24" + +[dependencies.url] +version = "2.5" + +[dev-dependencies.criterion] +version = "0.5" +features = ["async_tokio"] + +[dev-dependencies.tokio] +version = "1.50.0" +features = ["test-util"] + +[target."cfg(windows)".dependencies.windows-sys] +version = "0.61.2" +features = [ + "Win32_Foundation", + "Win32_System_Threading", +] diff --git a/vendor/mentra/Cargo.toml.orig b/vendor/mentra/Cargo.toml.orig new file mode 100644 index 0000000..38c1b6b --- /dev/null +++ b/vendor/mentra/Cargo.toml.orig @@ -0,0 +1,68 @@ +[package] +name = "mentra" +version = "0.18.3" +edition.workspace = true +rust-version.workspace = true +description = "An agent runtime for tool-using LLM applications" +license.workspace = true +repository.workspace = true +homepage.workspace = true +documentation = "https://docs.rs/mentra" +readme = "README.md" +keywords = ["agents", "llm", "runtime", "tools", "ai"] +categories = ["asynchronous", "development-tools"] + +[features] +# Default-on so an existing dependant sees no change: every caller who linked +# mentra before this feature existed had the Responses websocket transport, and +# an upgrade should not quietly take a transport away. Forwarded rather than +# left to mentra-provider's own default (which is why the dependency below sets +# `default-features = false`) so a host that only ever streams over HTTP+SSE can +# turn it off here and drop tokio-tungstenite from its tree. +default = ["responses-websocket"] +responses-websocket = ["mentra-provider/responses-websocket"] +openai-oauth = ["dep:ring"] +test-utils = [] + +[dependencies] +mentra-provider = { version = "0.5.1", path = "../mentra-provider", default-features = false } +tokio = { version = "1.50.0", features = ["full"] } +reqwest = { version = "0.12.23", default-features = false, features = [ + "json", + "rustls-tls", + "stream", +] } +url = "2.5" +serde = { version = "1.0.228", features = ["derive"] } +serde_json = "1.0.149" +async-trait = "0.1.89" +futures-util = "0.3.31" +time = { version = "0.3", features = ["formatting", "parsing", "serde"] } +base64 = "0.22.1" +serde_yaml_ng = "0.10" +directories = "6.0.0" +rusqlite = { version = "0.39", features = ["bundled"] } +thiserror = "2.0.18" +strum = { version = "0.27", features = ["derive"] } +libc = "0.2" +regex = "1.12.2" +rand = "0.9.2" +ring = { version = "0.17.14", optional = true } +glob-match = { workspace = true } +similar = "2.7.0" +unicode-normalization = "0.1.24" + +[dev-dependencies] +criterion = { version = "0.5", features = ["async_tokio"] } +# Tests only: `start_paused` lets a retry schedule measured in tens of seconds +# be asserted without waiting them out. Not enabled for the library build, so +# nothing a dependant links changes. +tokio = { version = "1.50.0", features = ["test-util"] } + +[target.'cfg(windows)'.dependencies] +windows-sys = { version = "0.61.2", features = ["Win32_Foundation", "Win32_System_Threading"] } + +[[bench]] +name = "long_session" +harness = false +required-features = ["test-utils"] diff --git a/vendor/mentra/LICENSE b/vendor/mentra/LICENSE new file mode 100644 index 0000000..62fd3c7 --- /dev/null +++ b/vendor/mentra/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 Wendell Wang + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/vendor/mentra/METACRATE.md b/vendor/mentra/METACRATE.md new file mode 100644 index 0000000..9131438 --- /dev/null +++ b/vendor/mentra/METACRATE.md @@ -0,0 +1,6 @@ +# MetaCrate adaptation + +This is Mentra 0.18.3 with a host-facing, bounded memory administration API. +`MemoryStore` adds stable list, agent-scoped detail, and count operations; +the volatile, SQLite, and hybrid stores implement the same contract. Storage +schema and tombstone behavior remain owned by Mentra. diff --git a/vendor/mentra/README.md b/vendor/mentra/README.md new file mode 100644 index 0000000..dc3b78f --- /dev/null +++ b/vendor/mentra/README.md @@ -0,0 +1,811 @@ +# mentra + +Mentra is an agent runtime for building tool-using LLM applications. + +MSRV: Rust 1.88. + +## Current Features + +- streaming model response handling +- provider-neutral token usage reporting across OpenAI, OpenRouter, Anthropic, Gemini, Ollama, and LM Studio +- optional tool authorization with structured previews and fail-closed execution blocking +- recoverable malformed tool-call input handling that feeds retry guidance back to the model +- custom tool execution through `ToolDefinition + ToolExecutor`, with `ToolSpec::builder(...)` as the convenience metadata API +- builtin `shell`, `background_run`, `check_background`, and `files` tools +- builtin `task` subagents with isolated child context and parent-side tracking +- persistent agent teams with `team_spawn`, `team_send`, `broadcast`, `team_read_inbox`, and generic request-response protocols via `team_request`, `team_respond`, and `team_list_requests` +- three-layer context compaction with silent tool-result shrinking, auto-summary compaction, and a builtin `compact` tool +- Model Context Protocol servers over stdio and the legacy HTTP+SSE transport, with their tools bridged into the runtime +- agent events and snapshots for CLI or UI watchers +- Anthropic provider support +- Gemini Developer API provider support +- OpenAI provider support via the Responses API +- OpenRouter provider support via the Responses API +- Ollama provider support via the OpenAI-compatible Responses API +- LM Studio provider support via the OpenAI-compatible Responses API +- image inputs for OpenAI and Anthropic, plus inline image bytes for Gemini + +## Quickstart Example + +Clone the repository and run the workspace quickstart example: + +```bash +cargo run -p mentra-examples --example quickstart -- "Summarize the benefits of tool-using agents." +``` + +The quickstart example accepts a prompt from CLI args or stdin. Set `MENTRA_MODEL` to force a specific OpenAI model; otherwise it resolves the newest available OpenAI model automatically. + +## Building A Runtime + +Use `Runtime::builder()` when you want Mentra's builtin runtime tools, or `Runtime::empty_builder()` when you want to opt into every tool explicitly. + +```rust,no_run +use mentra::{BuiltinProvider, Runtime}; + +fn main() -> Result<(), Box> { + let runtime = Runtime::builder() + .with_provider(BuiltinProvider::OpenAI, std::env::var("OPENAI_API_KEY")?) + .with_optional_provider( + BuiltinProvider::OpenRouter, + std::env::var("OPENROUTER_API_KEY").ok(), + ) + .with_optional_provider( + BuiltinProvider::Gemini, + std::env::var("GEMINI_API_KEY").ok(), + ) + .with_ollama() + .with_lmstudio() + .build()?; + + let _ = runtime; + Ok(()) +} +``` + +`with_ollama()` targets `http://127.0.0.1:11434/` and `with_lmstudio()` targets +`http://127.0.0.1:1234/`, using each server's OpenAI-compatible API surface. + +## Custom Compatible Providers + +If you need a non-default OpenAI-compatible or Anthropic-compatible endpoint, +register a provider-core instance with a customized `ProviderDefinition`. +Using a distinct provider ID lets you keep the builtin provider alongside your +custom endpoint. + +```rust,no_run +use mentra::{ModelSelector, ProviderId, Runtime}; + +# async fn demo() -> Result<(), Box> { +let mut definition = mentra::provider_core::responses::openai_definition(); +definition.descriptor.id = ProviderId::new("custom-openai-compatible"); +definition.descriptor.display_name = Some("Custom OpenAI-Compatible".to_string()); +definition.base_url = Some("https://llm.example.com/".to_string()); + +let runtime = Runtime::builder() + .with_registered_provider(mentra::provider_core::responses::ResponsesProvider::new( + definition, + mentra::provider_core::StaticCredentialSource::new(std::env::var("CUSTOM_API_KEY")?), + )) + .build()?; + +let model = runtime + .resolve_model( + ProviderId::new("custom-openai-compatible"), + ModelSelector::NewestAvailable, + ) + .await?; +# let _ = model; +# Ok(()) +# } +``` + +Anthropic-compatible endpoints follow the same pattern: + +```rust,no_run +use mentra::{ProviderId, Runtime}; + +# fn demo() -> Result<(), Box> { +let mut definition = mentra::provider_core::anthropic::definition(); +definition.descriptor.id = ProviderId::new("custom-anthropic-compatible"); +definition.descriptor.display_name = Some("Custom Anthropic-Compatible".to_string()); +definition.base_url = Some("https://claude.example.com/".to_string()); + +let runtime = Runtime::builder() + .with_registered_provider( + mentra::provider_core::anthropic::AnthropicProvider::with_definition_and_credential_source( + definition, + mentra::provider_core::StaticCredentialSource::new(std::env::var("CUSTOM_API_KEY")?), + ), + ) + .build()?; +# let _ = runtime; +# Ok(()) +# } +``` + +If your compatible endpoint needs different auth or extra headers, mutate the +definition's `auth_scheme`, `headers`, `query_params`, or `retry` fields before +registering it. + +## Architecture + +Mentra is organized around four runtime subsystems: + +- execution: model providers, runtime policy, hooks, turn execution, and shell/background command routing +- persistence: agent records, run state, task snapshots, leases, team state, background notifications, and memory +- tooling: builtin and custom tools, optional skills, and typed app context +- collaboration: persistent teammates, team inbox/request flows, and background task wakeups + +Persistent teammates are hosted as async actors on a shared Tokio runtime. Live actors are wake-driven rather than steady-state polled: inbox appends, protocol updates, background task completion, explicit resume, and autonomy timers wake the actor to process durable state already written to the store. After a restart, the persisted team inbox, protocol requests, and background notifications remain the source of truth, and `Runtime::resume(...)` revives teammate actors against that stored state. + +## Resolving A Model + +Use `Runtime::resolve_model(...)` when you want provider-aware model selection without reimplementing discovery or `ModelInfo` construction in application code. + +```rust,no_run +use mentra::{BuiltinProvider, ModelSelector, Runtime}; + +# async fn demo() -> Result<(), Box> { +let runtime = Runtime::builder() + .with_provider(BuiltinProvider::OpenAI, std::env::var("OPENAI_API_KEY")?) + .build()?; +let model = runtime + .resolve_model( + BuiltinProvider::OpenAI, + std::env::var("MENTRA_MODEL") + .map(ModelSelector::Id) + .unwrap_or(ModelSelector::NewestAvailable), + ) + .await?; + +let _ = model; +# Ok(()) +# } +``` + +## Coding Agent Setup + +`Runtime::builder()` registers Mentra's builtin tools, including `shell`, `background_run`, `check_background`, `files`, and the runtime/task/team intrinsics. Shell and background execution remain disabled by default, so coding-agent setups must opt in with a runtime policy. If you want semantic review before tools execute, install a `ToolAuthorizer`. + +The builtin local executor is a host executor, not a filesystem or network +sandbox. `RuntimePolicy::permissive()` therefore grants the model the same host +access as the Mentra process. Use it only inside a disposable container or +another boundary you trust. On a normal host, install an OS-enforced custom +executor with `RuntimeBuilder::with_executor(...)`; authorization and shell +validation decide whether a command may start, but they do not contain an +allowed command. + +For Responses API transport, xipe-compatible endpoints, and provider-side state +options, see the workspace +[`Responses Coding Agent Guide`](../docs/responses-coding-agent.md). + +```rust,no_run +use async_trait::async_trait; +use mentra::{BuiltinProvider, Runtime, RuntimePolicy}; +use mentra::tool::{ + ToolAuthorizationDecision, ToolAuthorizationRequest, ToolAuthorizer, +}; + +struct AllowAllAuthorizer; + +#[async_trait] +impl ToolAuthorizer for AllowAllAuthorizer { + async fn authorize( + &self, + _request: &ToolAuthorizationRequest, + ) -> Result { + Ok(ToolAuthorizationDecision::allow()) + } +} + +fn main() -> Result<(), Box> { + let runtime = Runtime::builder() + .with_provider(BuiltinProvider::OpenAI, std::env::var("OPENAI_API_KEY")?) + // Full host shell access. Use only inside a trusted external sandbox. + .with_policy(RuntimePolicy::permissive()) + .with_tool_authorizer(AllowAllAuthorizer) + .build()?; + + let _ = runtime; + Ok(()) +} +``` + +## Runtime Policy Defaults + +Mentra's builtin runtime tools are available by default, but command execution is not: + +- `Runtime::builder()` registers the builtin shell, background, file, task, team, and memory-oriented intrinsics +- foreground shell execution is disabled by default +- background command execution is disabled by default +- `RuntimePolicy::permissive()` enables both shell and background command execution +- `RuntimePolicy::workspace_bounded(...)` and `RuntimePolicy::read_only(...)` keep shell execution disabled; their roots constrain builtin file tools and the requested shell working directory, not shell process effects +- builtin shell commands run through `/bin/sh -c` on Unix and `cmd.exe /C` on Windows +- the local executor clears unlisted environment variables and enforces timeouts, output caps, and process-tree cleanup on timeout, but it does not restrict filesystem or network access +- semantic review is opt-in through `RuntimeBuilder::with_tool_authorizer(...)` + +Use the default policy when you want a safer runtime surface. Opt into +`RuntimePolicy::permissive()` only when an external sandbox already contains the +entire Mentra process and full host access is intentional. + +If you need different command semantics, such as PowerShell on Windows, or +filesystem/network confinement, replace the default local executor with +`RuntimeBuilder::with_executor(...)`. A workspace-bounded or read-only policy +can then explicitly enable foreground and background shell switches; Mentra +treats that executor as a trusted enforcement boundary and does not fall back +to the local executor. + +## Tool Authorization + +Mentra can run a caller-provided authorization pass before any tool executes. This is the recommended integration point for LLM-based security review, human approval, or custom policy engines. + +- no authorizer installed: tools run under the remaining hard runtime constraints +- authorizer returns `Allow`: the tool executes +- authorizer returns `Prompt` or `Deny`: Mentra blocks execution and returns an error `tool_result` +- authorizer timeout or error: Mentra fails closed and blocks execution + +Every authorization request includes a `ToolAuthorizationPreview` with tool metadata plus structured input. Builtin tools provide more specific previews: + +- `shell` and `background_run` include the raw command, resolved working directory, timeout, background flag, and justification +- `files` includes resolved paths and operation kinds such as `read`, `search`, `set`, `move`, and `delete`, without file contents + +```rust,no_run +use async_trait::async_trait; +use mentra::tool::{ + ToolAuthorizationDecision, ToolAuthorizationRequest, ToolAuthorizer, +}; + +struct DenyDeletes; + +#[async_trait] +impl ToolAuthorizer for DenyDeletes { + async fn authorize( + &self, + request: &ToolAuthorizationRequest, + ) -> Result { + let structured = &request.preview.structured_input; + let denies_delete = structured + .get("operations") + .and_then(|value| value.as_array()) + .is_some_and(|ops| ops.iter().any(|op| op.get("op").and_then(|v| v.as_str()) == Some("delete"))); + + if request.tool_name == "files" && denies_delete { + Ok(ToolAuthorizationDecision::deny("delete operations require manual approval")) + } else { + Ok(ToolAuthorizationDecision::allow()) + } + } +} +``` + +Registering a skills directory also makes the builtin `load_skill` tool available: + +```rust,no_run +use mentra::{BuiltinProvider, Runtime}; + +fn main() -> Result<(), Box> { + let runtime = Runtime::builder() + .with_provider(BuiltinProvider::OpenAI, std::env::var("OPENAI_API_KEY")?) + .with_skills_dir("./skills")? + .build()?; + + let _ = runtime; + Ok(()) +} +``` + +## App Context + +If your tools need access to typed host-side state, register it on the runtime and retrieve it from `ToolContext` or `ParallelToolContext`: + +```rust,no_run +use std::sync::Arc; + +use async_trait::async_trait; +use mentra::{ + BuiltinProvider, Runtime, + tool::{ToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSpec}, +}; +use serde_json::{Value, json}; + +struct AppState { + api_base: String, +} + +struct InspectStateTool; + +impl ToolDefinition for InspectStateTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("inspect_state") + .description("Return the configured API base URL.") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .build() + } +} + +#[async_trait] +impl ToolExecutor for InspectStateTool { + async fn execute_mut(&self, ctx: ToolContext<'_>, _input: Value) -> ToolResult { + let state = ctx.app_context::()?; + Ok(state.api_base.clone()) + } +} + +fn main() -> Result<(), Box> { + let runtime = Runtime::builder() + .with_provider(BuiltinProvider::OpenAI, std::env::var("OPENAI_API_KEY")?) + .with_context(Arc::new(AppState { + api_base: "https://api.example.com".to_string(), + })) + .with_tool(InspectStateTool) + .build()?; + + let _ = runtime; + Ok(()) +} +``` + +## Custom Tools + +Use `ToolSpec::builder(...)` to define custom tools without hand-assembling the metadata struct: + +```rust,no_run +use async_trait::async_trait; +use mentra::tool::{ + ParallelToolContext, ToolCapability, ToolDefinition, ToolDurability, ToolExecutor, + ToolResult, ToolSideEffectLevel, ToolSpec, +}; +use serde_json::{Value, json}; + +struct UppercaseTool; + +impl ToolDefinition for UppercaseTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("uppercase_text") + .description("Uppercase the provided text") + .input_schema(json!({ + "type": "object", + "properties": { + "text": { "type": "string" } + }, + "required": ["text"] + })) + .capability(ToolCapability::ReadOnly) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_timeout(std::time::Duration::from_secs(5)) + .build() + } +} + +#[async_trait] +impl ToolExecutor for UppercaseTool { + async fn execute(&self, _ctx: ParallelToolContext, input: Value) -> ToolResult { + let text = input + .get("text") + .and_then(|value| value.as_str()) + .ok_or_else(|| "text is required".to_string())?; + Ok(text.to_uppercase()) + } +} +``` + +`ToolSpec::execution_timeout(...)` is enforced by Mentra around the tool future itself, which is useful for network-backed tools that need a tighter budget than the overall agent run. + +Internally, Mentra translates `ToolSpec` into a runtime-only `RuntimeToolDescriptor`, but custom runtime integrations should continue to treat `ToolSpec::builder(...)` as the supported public metadata surface. `ExecutableTool` remains available in this release as a compatibility trait alias over `ToolDefinition + ToolExecutor`. + +When a tool needs disposable delegated work, `ParallelToolContext::spawn_subagent()` can create a child agent that inherits the current runtime and model defaults. See the `subagent_tool` example in the workspace examples crate for a complete usage pattern. + +Override `ToolExecutor::authorization_preview(...)` when your custom tool needs to expose structured metadata to the installed `ToolAuthorizer`. The default preview includes the resolved working directory, tool capabilities, side-effect level, durability, the raw JSON input, and the same JSON as `structured_input`. + +## Tooling Layers + +Mentra now separates tool contracts into explicit layers: + +- `ProviderToolSpec` in `mentra-provider` for provider-facing serialization +- `RuntimeToolDescriptor` in Mentra for scheduling, approval, and durability metadata +- `ToolDefinition + ToolExecutor` for executable runtime tools + +Provider adapters should serialize provider-facing tool specs only. Runtime integrations should continue to implement custom tools with `ToolSpec::builder(...)`, `ToolDefinition`, and `ToolExecutor`. + +## Hosted Tool Search + +Mentra can mark custom tools as deferred and let a provider load them on demand with native hosted tool search. + +Mark a tool as deferred in its `ToolSpec`: + +```rust,no_run +use async_trait::async_trait; +use mentra::tool::{ParallelToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSpec}; +use serde_json::{Value, json}; + +struct LookupOrderTool; + +impl ToolDefinition for LookupOrderTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("lookup_order") + .description("Look up an order by id.") + .input_schema(json!({ + "type": "object", + "properties": { + "order_id": { "type": "string" } + }, + "required": ["order_id"] + })) + .defer_loading(true) + .build() + } +} + +#[async_trait] +impl ToolExecutor for LookupOrderTool { + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + Ok("order loaded".to_string()) + } +} +``` + +Enable hosted tool search per agent with `ProviderRequestOptions`: + +```rust,no_run +use mentra::agent::AgentConfig; +use mentra::provider::{ProviderRequestOptions, ReasoningEffort, ReasoningOptions, ToolSearchMode}; + +let config = AgentConfig { + provider_request_options: ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Hosted, + reasoning: Some(ReasoningOptions { + effort: Some(ReasoningEffort::Medium), + summary: None, + }), + ..Default::default() + }, + ..Default::default() +}; +``` + +Current provider support: + +- OpenAI: supported through the Responses API hosted `tool_search` surface +- Anthropic: supported through the Messages API BM25 tool-search server tool +- Gemini: deferred custom tools are not supported; Mentra returns `InvalidRequest` + +Reasoning effort support: + +- The shared levels are `low`, `medium`, `high`, `xhigh`, and `max`; omitting + effort leaves the provider default unchanged. +- OpenAI and OpenRouter: Mentra forwards all five levels as + `reasoning.effort` on the Responses API. +- Anthropic: Mentra writes the requested level to `output_config.effort` and + enables adaptive thinking on models that support it. Opus 4.5 accepts + `low`/`medium`/`high` effort without adaptive thinking; availability of + `xhigh` and `max` depends on the Claude model. +- Gemini: Mentra maps the shared `low`, `medium`, and `high` levels to + `thinkingLevel` on Gemini 3 models, subject to that model's accepted values. + `xhigh` and `max` return `InvalidRequest` instead of being silently + downgraded. +- Anthropic models without effort support and Gemini models older than 3 return + `InvalidRequest` when unified reasoning effort is set. + +Deferred tools are filtered through `ToolProfile` just like immediate tools. If you force a deferred tool with `ToolChoice::Tool { name }`, Mentra serializes that specific tool as immediate for the request so explicit invocation still works. + +## Model Context Protocol Servers + +Mentra connects to external MCP servers and bridges every tool they advertise +into the runtime under a namespaced `mcp____` name. Bridged tools +run through the same authorization, result limiter, and paging path as builtin +and custom tools. + +Two transports are supported, selected by which configuration type you register. + +**stdio** spawns the server as a child process: + +```rust,no_run +use mentra::{BuiltinProvider, McpServerConfig, Runtime}; + +# async fn demo() -> Result<(), Box> { +let runtime = Runtime::builder() + .with_provider(BuiltinProvider::Anthropic, std::env::var("ANTHROPIC_API_KEY")?) + .with_mcp_server(McpServerConfig { + name: "filesystem".to_string(), + command: "npx".to_string(), + args: vec![ + "-y".to_string(), + "@modelcontextprotocol/server-filesystem".to_string(), + "/tmp".to_string(), + ], + env: Default::default(), + cwd: None, + }) + .build_async() + .await?; +# let _ = runtime; +# Ok(()) +# } +``` + +**Legacy HTTP+SSE** reaches a hosted server over the network: + +```rust,no_run +use mentra::{BuiltinProvider, McpSseServerConfig, Runtime}; + +# async fn demo() -> Result<(), Box> { +let runtime = Runtime::builder() + .with_provider(BuiltinProvider::Anthropic, std::env::var("ANTHROPIC_API_KEY")?) + .with_mcp_sse_server( + McpSseServerConfig::new("observability", "https://mcp.example.com/sse") + .with_bearer_token(std::env::var("MCP_TOKEN")?), + ) + .build_async() + .await?; +# let _ = runtime; +# Ok(()) +# } +``` + +A server that answers `404` on `/mcp` but serves `/sse` needs this transport. + +### HTTP+SSE is not Streamable HTTP + +`McpSseServerConfig` speaks the transport from MCP protocol revision +`2024-11-05`, which is a different protocol from the newer Streamable HTTP: + +| | legacy HTTP+SSE | Streamable HTTP | +|---|---|---| +| Endpoints | a `GET` stream plus a separate `POST` URL | one URL for both | +| POST target | named by the server in an `endpoint` event | the configured URL | +| Responses | always on the `GET` stream | in the POST response or a stream | +| Session | a query parameter in the endpoint URL | the `Mcp-Session-Id` header | + +The client opens the configured URL with `Accept: text/event-stream`, waits for +an `endpoint` event naming the POST URL, then posts `initialize`, a +`notifications/initialized` notification, and a paginated `tools/list`. Servers +answer each POST `202 Accepted` and deliver the actual JSON-RPC result as a +`message` event on the stream. + +### Security and failure behavior + +The endpoint URL is chosen by the server, so it is validated before anything is +sent to it. A resolved endpoint must match the configured URL's scheme, host, +and effective port; a cross-origin endpoint, a protocol-relative `//other.host` +value, embedded credentials, and non-`http(s)` schemes are all refused. Redirects +are never followed on either request. + +Configured headers are sent on both the stream and every POST, stored as +`SecretString` so they never appear in `Debug` output, errors, or logs. +Configuring headers against a plaintext `http://` URL on a non-loopback host is +rejected unless `allowing_plaintext_credentials()` is set. No error carries a +response body or SSE payload, so a malicious server cannot write text into your +logs. + +Losing the stream ends the session — the client fails closed rather than +hanging, and never reconnects or re-sends a `tools/call`. A call whose response +never arrived surfaces as `McpSseError::RequestIndeterminate`, because the POST +and the response travel on different connections: the tool may have run. Treat +that differently from a rejected POST, which definitely did not execute. + +### Using the client directly + +Hosts that need their own allowlist, redaction, or evidence policy can drive +`McpSseClient` without registering anything: + +```rust,no_run +use mentra::{McpSseClient, McpSseServerConfig}; + +# async fn demo() -> Result<(), Box> { +let config = McpSseServerConfig::new("observability", "https://mcp.example.com/sse") + .with_bearer_token(std::env::var("MCP_TOKEN")?); + +let client = McpSseClient::connect(&config).await?; +for tool in client.tools() { + println!("{}", tool.name); +} + +let result = client + .call_tool("search_logs", Some(serde_json::json!({"query": "error"}))) + .await?; +println!("{}", result.is_error); + +client.shutdown().await; +# Ok(()) +# } +``` + +## Tool Profiles + +Register tools once on the runtime, then use `AgentConfig::tool_profile` to expose different subsets for different operating modes. + +```rust,no_run +use mentra::{BuiltinProvider, ModelSelector, Runtime}; +use mentra::agent::{AgentConfig, ToolProfile}; + +# async fn demo() -> Result<(), Box> { +let runtime = Runtime::builder() + .with_provider(BuiltinProvider::OpenAI, std::env::var("OPENAI_API_KEY")?) + .build()?; +let model = runtime + .resolve_model( + BuiltinProvider::OpenAI, + ModelSelector::Id("gpt-5.4-mini".to_string()), + ) + .await?; + +let queue_mode = AgentConfig { + tool_profile: ToolProfile::only([ + "shell", + "background_run", + "check_background", + "files", + "task", + ]), + ..Default::default() +}; + +let direct_mode = AgentConfig { + tool_profile: ToolProfile::hide(["task", "background_run"]), + ..Default::default() +}; + +let _queue_agent = runtime.spawn_with_config("Queue Agent", model.clone(), queue_mode)?; +let _direct_agent = runtime.spawn_with_config("Direct Agent", model, direct_mode)?; +# Ok(()) +# } +``` + +This is the recommended pattern when one application needs multiple tool surfaces such as a queue-backed agent with delegation enabled and a direct mode that keeps the same runtime but hides long-running or task-oriented tools. + +## CLI Integration Pattern + +For CLI-style coding or analysis tools, the usual setup is: + +- register a superset of builtin and custom tools on one runtime +- scope shell and file access with `RuntimePolicy` +- keep application-specific output paths in app context for custom tools +- switch behavior per mode by changing `AgentConfig::tool_profile`, not by rebuilding the runtime +- inspect `agent.history()` after the run when you want to render a compact tool log or transcript summary + +The `cli_runtime` example in the workspace examples crate shows this pattern end to end with custom tools, policy setup, mode-specific tool surfaces, and transcript inspection. + +## Disposable Tasks vs Persistent Teams + +Mentra supports two different delegation models: + +- use the builtin `task` tool or `ParallelToolContext::spawn_subagent()` for short-lived disposable delegation that should return a single summary to the parent +- use `team_spawn`, `team_send`, `team_read_inbox`, `team_request`, and `team_respond` when you want a persistent teammate with a durable mailbox and request/response workflow across turns + +The `task` path is ideal for one-off decomposition inside a single run. The `team_*` tools are for longer-lived collaborators that should keep state, receive follow-up work, and participate in approval or shutdown flows. + +## Sending Images + +You can attach image blocks alongside text when sending a user turn: + +```rust,no_run +# use mentra::{ContentBlock, Agent}; +# async fn demo(agent: &mut Agent) -> Result<(), Box> { +agent + .send(vec![ + ContentBlock::text("What is happening in this screenshot?"), + ContentBlock::image_bytes("image/png", std::fs::read("screenshot.png")?), + ]) + .await?; +# Ok(()) +# } +``` + +For already-hosted assets, use `ContentBlock::image_url(...)` instead. Gemini currently supports inline `image_bytes(...)` inputs only and rejects `image_url(...)`. + +## Long-Term Memory + +Agents automatically recall from long-term memory by default. When you use `Runtime::builder()`, the builtin runtime intrinsics include: + +- `memory_search` for explicit recall +- `memory_pin` for writing important facts +- `memory_forget` for tombstoning a specific memory record + +`MemoryConfig` controls recall and write behavior per agent. The default configuration enables automatic recall and memory write tools, which is useful for long-running assistants and teammate workflows. Disable write tools when you want recall without model-initiated mutation. + +## Context Compaction + +Agents compact context by default: + +- old tool results are micro-compacted in outbound requests +- when estimated request context exceeds roughly 50k tokens, Mentra writes the full transcript to the default transcript directory and replaces older history with a model-generated summary +- the model can also call the builtin `compact` tool explicitly + +You can tune or disable this per-agent with `CompactionConfig`: + +```rust +use mentra::agent::{AgentConfig, CompactionConfig}; + +let config = AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(75_000), + ..Default::default() + }, + ..Default::default() +}; +``` + +## Data And Persistence Defaults + +For non-test builds, Mentra keeps all default persisted state under a workspace-scoped app-data directory: + +- store: `/mentra/workspaces//runtime.sqlite` +- runtime-scoped stores: `/mentra/workspaces//runtime-.sqlite` +- team state: `/mentra/workspaces//team/` +- task state: `/mentra/workspaces//tasks/` +- transcripts: `/mentra/workspaces//transcripts/` + +If the platform data directory cannot be resolved, Mentra falls back to `.mentra/workspaces//...` inside the current workspace. + +Override these defaults when needed: + +- use `Runtime::builder().with_store(...)` for the SQLite store +- customize `AgentConfig::task.tasks_dir`, `AgentConfig::team.team_dir`, and `AgentConfig::compaction.transcript_dir` for task, team, and transcript storage + +## Persistence Extension Points + +The public persistence surface is intentionally split into narrower traits: + +- `AgentStore` for agent records and working-memory snapshots +- `RunStore` for turn and run lifecycle tracking +- `TaskStore` for the dependency-aware task board +- `LeaseStore` for runtime ownership and resume coordination + +`RuntimeStore` composes those traits with `TeamStore`, `BackgroundStore`, and `MemoryStore`. `SqliteRuntimeStore` is the default all-in-one backend. `HybridRuntimeStore` keeps SQLite runtime state and swaps in the hybrid memory engine for richer long-term memory behavior. + +## Testing With MockRuntime + +Enable the `test-utils` feature when you want a deterministic scripted runtime for unit and integration tests. + +`mentra::test::MockRuntime` wraps a real runtime with: + +- a scripted provider +- a `VolatileRuntimeStore`, so a mock writes nothing to disk and two mocks never + share state — pass `MockRuntimeBuilder::with_store` a `SqliteRuntimeStore` + when a test needs state that outlives the mock +- deterministic per-turn helper methods for assistant text, streamed text, tool-call turns, and provider failures + +This is the recommended way to test Mentra-based agents and tools without live API keys. + +The common pattern is: + +- build a `MockRuntime` +- register the same custom tools you use in production +- spawn an agent with the `AgentConfig` or `ToolProfile` you want to verify +- assert against `mock.recorded_requests()` to confirm the runtime exposed the expected tools and tool-choice hints + +See `mentra::test` and the crate tests for a full example of asserting runtime assembly with custom tools and filtered tool surfaces. + +## Interactive Repo Example + +Clone the repository when you want the richer interactive demo with provider selection, persisted runtime inspection, skills loading, and team/task visibility. + +Set `OPENAI_API_KEY`, `OPENROUTER_API_KEY`, `ANTHROPIC_API_KEY`, or `GEMINI_API_KEY`, then run. The example lets you choose a provider and shows up to 10 models from that provider ordered newest to oldest. + +```bash +cargo run -p mentra-examples --example chat +``` + +Additional focused examples live in the same crate: + +```bash +cargo run -p mentra-examples --example custom_tool +cargo run -p mentra-examples --example subagent_tool +cargo run -p mentra-examples --example team_collaboration +cargo run -p mentra-examples --example cli_runtime -- --mode direct +``` + +`cli_runtime` is the closest example to a real integration. It combines runtime policy setup, custom tools, mode-specific `ToolProfile` selection, and transcript inspection after the run. + +## Run Checks + +```bash +cargo fmt --all --check +cargo check --workspace +cargo clippy --workspace --all-targets -- -D warnings +cargo test --workspace +``` diff --git a/vendor/mentra/benches/long_session.rs b/vendor/mentra/benches/long_session.rs new file mode 100644 index 0000000..ed9afdd --- /dev/null +++ b/vendor/mentra/benches/long_session.rs @@ -0,0 +1,186 @@ +use criterion::{Criterion, criterion_group, criterion_main}; +use mentra::{ + ContentBlock, + memory::{MemoryRecord, MemoryRecordKind, MemorySearchMode, MemorySearchRequest, MemoryStore}, + runtime::SqliteRuntimeStore, + test::{MockRuntimeBuilder, MockTurn}, +}; +use std::time::SystemTime; + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +fn now_secs() -> i64 { + SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64 +} + +fn temp_sqlite_path(label: &str) -> std::path::PathBuf { + std::env::temp_dir().join(format!( + "mentra-bench-{}-{}.sqlite", + label, + SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() + )) +} + +fn build_mock_runtime( + n_turns: usize, + store_path: std::path::PathBuf, +) -> (mentra::test::MockRuntime, mentra::ModelInfo) { + let store = SqliteRuntimeStore::new(store_path); + let mut builder = MockRuntimeBuilder::default().with_store(store); + for i in 0..n_turns { + builder = builder.push_turn(MockTurn::Text(format!("turn {i} response"))); + } + let mock = builder.build().expect("build mock runtime"); + let model = mock.model(); + (mock, model) +} + +// --------------------------------------------------------------------------- +// bench_500_turn_session +// Measures total wall time for driving 500 send/response cycles. +// --------------------------------------------------------------------------- + +fn bench_500_turn_session(c: &mut Criterion) { + let rt = tokio::runtime::Runtime::new().expect("tokio runtime"); + + c.bench_function("500_turn_session", |b| { + b.to_async(&rt).iter(|| async { + let store_path = temp_sqlite_path("500turns"); + let (mock, model) = build_mock_runtime(500, store_path); + let mut agent = mock + .runtime() + .spawn("bench-agent", model) + .expect("spawn agent"); + + for i in 0u32..500 { + agent + .send(vec![ContentBlock::text(format!("message {i}"))]) + .await + .expect("send turn"); + } + }); + }); +} + +// --------------------------------------------------------------------------- +// bench_resume_after_heavy_session +// 200 turns followed by a resume from the persisted agent record. +// --------------------------------------------------------------------------- + +fn bench_resume_after_heavy_session(c: &mut Criterion) { + let rt = tokio::runtime::Runtime::new().expect("tokio runtime"); + + c.bench_function("resume_after_200_turn_session", |b| { + b.to_async(&rt).iter(|| async { + let store_path = temp_sqlite_path("resume"); + // Prepare: 200 turns for the heavy session + 1 turn for the resumed send + let store = SqliteRuntimeStore::new(store_path.clone()); + let mut builder = MockRuntimeBuilder::default().with_store(store); + for i in 0..201usize { + builder = builder.push_turn(MockTurn::Text(format!("turn {i} response"))); + } + let mock = builder.build().expect("build mock runtime"); + let model = mock.model(); + + // Drive the heavy session + let agent_id = { + let mut agent = mock + .runtime() + .spawn("bench-agent", model.clone()) + .expect("spawn agent"); + for i in 0u32..200 { + agent + .send(vec![ContentBlock::text(format!("message {i}"))]) + .await + .expect("send turn"); + } + agent.id().to_string() + }; + + // Measure: resume and send one more turn + let mut resumed = mock + .runtime() + .resume_agent(&agent_id) + .expect("resume agent"); + resumed + .send(vec![ContentBlock::text("resumed message")]) + .await + .expect("send after resume"); + }); + }); +} + +// --------------------------------------------------------------------------- +// bench_memory_scaling_1000_records +// Seeds 1000 memory records into a SQLite store and measures search latency. +// --------------------------------------------------------------------------- + +fn bench_memory_scaling_1000_records(c: &mut Criterion) { + use mentra::memory::SqliteHybridMemoryStore; + + let rt = tokio::runtime::Runtime::new().expect("tokio runtime"); + + c.bench_function("memory_search_1000_records", |b| { + b.to_async(&rt).iter(|| async { + let store_path = temp_sqlite_path("memory"); + let store = SqliteHybridMemoryStore::new(&store_path); + + // Seed 1000 records + let records: Vec = (0..1000usize) + .map(|i| MemoryRecord { + record_id: format!("rec-{i:04}"), + agent_id: "bench-agent".to_string(), + kind: if i % 3 == 0 { + MemoryRecordKind::Episode + } else if i % 3 == 1 { + MemoryRecordKind::Fact + } else { + MemoryRecordKind::Summary + }, + content: format!( + "memory record {i}: the agent discussed topic {i} in session {i}" + ), + source_revision: i as u64, + created_at: now_secs() - i as i64, + metadata_json: "{}".to_string(), + source: None, + pinned: false, + score: None, + }) + .collect(); + + store.upsert_records(&records).expect("seed records"); + + // Measure search latency + let request = MemorySearchRequest { + agent_id: "bench-agent".to_string(), + query: "agent discussed topic session".to_string(), + limit: 20, + char_budget: None, + mode: MemorySearchMode::Automatic, + }; + let _hits = store + .search_records_with_options(&request) + .expect("search records"); + }); + }); +} + +// --------------------------------------------------------------------------- +// Criterion group +// --------------------------------------------------------------------------- + +criterion_group! { + name = benches; + config = Criterion::default().sample_size(10); + targets = bench_500_turn_session, bench_resume_after_heavy_session, bench_memory_scaling_1000_records +} +criterion_main!(benches); diff --git a/vendor/mentra/src/agent.rs b/vendor/mentra/src/agent.rs new file mode 100644 index 0000000..9e49ea4 --- /dev/null +++ b/vendor/mentra/src/agent.rs @@ -0,0 +1,648 @@ +mod compact; +mod config; +mod events; +mod lifecycle; +mod pending; +mod pending_block; +mod round_strategy; +mod runner; +mod snapshot; +mod steering; +mod subagent; +mod task_state; +mod team; +mod terminal_output; +#[cfg(test)] +mod tests; +mod wait; + +use std::{ + collections::HashSet, + sync::{ + Arc, Mutex, + atomic::{AtomicU64, Ordering}, + }, +}; + +use serde::{Deserialize, Serialize}; +use tokio::sync::{broadcast, watch}; + +use crate::{ + ContentBlock, Message, + background::BackgroundNotification, + error::RuntimeError, + memory::journal::{AgentMemory, AgentMemoryState as MemoryState}, + provider::{Provider, ProviderId, ToolChoice}, + runtime::{ + LoadedAgentState, RuntimeIntrinsicTool, TaskItem, + handle::{AgentExecutionConfig, AgentObserver, RuntimeHandle}, + }, + team::TeamMessage, + transcript::{DelegationArtifact, DelegationEdge, TranscriptItem}, +}; + +pub(crate) use team::parse_task_input; + +pub use config::{ + AgentConfig, CompactionConfig, ContextCompactionConfig, MemoryConfig, TaskConfig, + TeamAutonomyConfig, TeamConfig, ToolProfile, ToolResultPagingConfig, WorkspaceConfig, +}; +pub use events::{ + AgentEvent, AgentSnapshot, AgentStatus, CompactionDetails, CompactionTrigger, + ContextCompactionDetails, ContextCompactionTrigger, PendingToolUseSummary, SpawnedAgentStatus, + SpawnedAgentSummary, +}; +pub use pending::PendingAssistantTurn; +pub use round_strategy::{ + ReasoningChange, RoundAdjustment, RoundBoundary, RoundContext, RoundDecision, RoundStrategy, + RoundToolResult, +}; +use runner::TurnRunner; +pub use steering::{QueueMode, SteeringHandle}; +pub(crate) use subagent::DisposableSubagentTemplate; +use terminal_output::TerminalToolGate; +pub use terminal_output::{FinalOutput, TerminalOutputSpec}; +pub use wait::{AgentWaitFuture, AgentWaitHandle}; + +static NEXT_AGENT_ID: AtomicU64 = AtomicU64::new(1); + +/// Running or persisted agent managed by a [`crate::Runtime`]. +pub struct Agent { + id: String, + runtime: RuntimeHandle, + model: String, + provider_id: ProviderId, + name: String, + config: AgentConfig, + memory: AgentMemory, + tasks: Vec, + rounds_since_task: usize, + event_bus: AgentEventBus, + snapshot: Arc>, + snapshot_tx: watch::Sender, + provider: Arc, + hidden_tools: HashSet, + terminal_tool_gate: Arc>>, + max_rounds: Option, + inflight_background_notifications: Vec, + inflight_team_messages: Vec, + steering: SteeringHandle, + inflight_steer: Vec>, + inflight_follow_up: Vec>, + teammate_identity: Option, + idle_requested: bool, + current_run_id: Option, + /// Full texts of results this agent received paged, keyed by + /// `tool_use_id` — the backing store for `read_tool_result`. Empty and + /// unused unless `config.tool_result_paging` is set. + paged_tool_results: crate::tool::paging::PagedToolResults, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub(crate) struct TeammateIdentity { + pub(crate) role: String, + pub(crate) lead: String, +} + +#[derive(Default)] +pub(crate) struct AgentSpawnOptions { + pub(crate) hidden_tools: HashSet, + pub(crate) max_rounds: Option, + pub(crate) teammate_identity: Option, +} + +type AgentEventTap = Arc; + +#[derive(Default)] +struct AgentEventTapRegistry { + next_id: u64, + taps: Vec<(u64, AgentEventTap)>, +} + +pub(crate) struct AgentEventTapGuard { + registry: Arc>, + id: u64, +} + +#[derive(Clone)] +pub(crate) struct AgentEventBus { + tx: broadcast::Sender, + taps: Arc>, +} + +impl AgentEventBus { + fn new(capacity: usize) -> Self { + let (tx, _) = broadcast::channel(capacity); + Self { + tx, + taps: Arc::new(Mutex::new(AgentEventTapRegistry::default())), + } + } + + pub(crate) fn send(&self, event: AgentEvent) { + let taps = { + let registry = self.taps.lock().expect("agent event tap registry poisoned"); + registry + .taps + .iter() + .map(|(_, tap)| Arc::clone(tap)) + .collect::>() + }; + for tap in taps { + tap(&event); + } + let _ = self.tx.send(event); + } + + pub(crate) fn subscribe(&self) -> broadcast::Receiver { + self.tx.subscribe() + } + + pub(crate) fn register_tap( + &self, + tap: impl Fn(&AgentEvent) + Send + Sync + 'static, + ) -> AgentEventTapGuard { + let mut registry = self.taps.lock().expect("agent event tap registry poisoned"); + let id = registry.next_id; + registry.next_id += 1; + registry.taps.push((id, Arc::new(tap))); + AgentEventTapGuard { + registry: Arc::clone(&self.taps), + id, + } + } +} + +impl Drop for AgentEventTapGuard { + fn drop(&mut self) { + let mut registry = self + .registry + .lock() + .expect("agent event tap registry poisoned"); + registry.taps.retain(|(tap_id, _)| *tap_id != self.id); + } +} + +impl Agent { + pub(crate) fn new( + runtime: RuntimeHandle, + model: String, + name: String, + config: AgentConfig, + provider: Arc, + options: AgentSpawnOptions, + ) -> Result { + let AgentSpawnOptions { + hidden_tools, + max_rounds, + teammate_identity, + } = options; + let store = runtime.store(); + let agent_id = format!( + "agent-{:x}-{}", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(), + NEXT_AGENT_ID.fetch_add(1, Ordering::Relaxed) + ); + let memory = AgentMemory::new(agent_id.clone(), store.clone(), MemoryState::default()); + let event_bus = AgentEventBus::new(256); + let memory_view = memory.snapshot_view(); + let snapshot = AgentSnapshot { + history_len: memory_view.history_len, + current_text: memory_view.current_text, + pending_tool_uses: memory_view.pending_tool_uses, + ..Default::default() + }; + let snapshot = Arc::new(Mutex::new(snapshot)); + let (snapshot_tx, _) = + watch::channel(snapshot.lock().expect("agent snapshot poisoned").clone()); + let mut agent = Self { + id: agent_id, + runtime, + model, + provider_id: provider.descriptor().id, + name, + config, + memory, + tasks: Vec::new(), + rounds_since_task: 0, + event_bus, + snapshot, + snapshot_tx, + provider, + hidden_tools, + terminal_tool_gate: Arc::new(Mutex::new(None)), + max_rounds, + inflight_background_notifications: Vec::new(), + inflight_team_messages: Vec::new(), + steering: SteeringHandle::new(), + inflight_steer: Vec::new(), + inflight_follow_up: Vec::new(), + teammate_identity, + idle_requested: false, + current_run_id: None, + paged_tool_results: Default::default(), + }; + agent + .runtime + .store() + .create_agent(&agent.persisted_record(), agent.memory.state())?; + let execution_config = AgentExecutionConfig { + name: agent.name.clone(), + team_dir: agent.config.team.team_dir.clone(), + tasks_dir: agent.config.task.tasks_dir.clone(), + base_dir: agent.config.workspace.base_dir.clone(), + memory_tool_search_limit: agent.config.memory.tool_search_limit, + auto_route_shell: agent.config.workspace.auto_route_shell, + is_teammate: agent.teammate_identity.is_some(), + }; + let observer = AgentObserver { + events: agent.event_bus.clone(), + snapshot_tx: agent.snapshot_tx.clone(), + snapshot: Arc::clone(&agent.snapshot), + }; + agent + .runtime + .register_agent(&agent.id, &agent.name, execution_config, &observer)?; + agent.register_tool_result_pager(); + agent.refresh_tasks_from_disk()?; + Ok(agent) + } + + pub(crate) fn from_loaded( + runtime: RuntimeHandle, + mut state: LoadedAgentState, + provider: Arc, + ) -> Result { + let mut memory = AgentMemory::new(state.record.id.clone(), runtime.store(), state.memory); + let recovery = memory.recover()?; + if recovery.interrupted { + state.record.status = AgentStatus::Interrupted; + runtime.store().update_run_state( + recovery + .interrupted_run_id + .as_deref() + .expect("recovery should include run id"), + "interrupted", + Some("recovered after interruption"), + )?; + runtime.store().save_agent_record(&state.record)?; + } + let memory_view = memory.snapshot_view(); + let snapshot = AgentSnapshot { + status: state.record.status.clone(), + history_len: memory_view.history_len, + current_text: memory_view.current_text, + pending_tool_uses: memory_view.pending_tool_uses, + pending_team_messages: 0, + subagents: state.record.subagents.clone(), + ..Default::default() + }; + let snapshot = Arc::new(Mutex::new(snapshot)); + let (snapshot_tx, _) = + watch::channel(snapshot.lock().expect("agent snapshot poisoned").clone()); + let event_bus = AgentEventBus::new(256); + let mut agent = Self { + id: state.record.id.clone(), + runtime, + model: state.record.model.clone(), + provider_id: state.record.provider_id.clone(), + name: state.record.name.clone(), + config: state.record.config.clone(), + memory, + tasks: Vec::new(), + rounds_since_task: state.record.rounds_since_task, + event_bus, + snapshot, + snapshot_tx, + provider, + hidden_tools: state.record.hidden_tools, + terminal_tool_gate: Arc::new(Mutex::new(None)), + max_rounds: state.record.max_rounds, + inflight_background_notifications: Vec::new(), + inflight_team_messages: Vec::new(), + steering: SteeringHandle::new(), + inflight_steer: Vec::new(), + inflight_follow_up: Vec::new(), + teammate_identity: state.record.teammate_identity, + idle_requested: state.record.idle_requested, + current_run_id: None, + paged_tool_results: Default::default(), + }; + let execution_config = AgentExecutionConfig { + name: agent.name.clone(), + team_dir: agent.config.team.team_dir.clone(), + tasks_dir: agent.config.task.tasks_dir.clone(), + base_dir: agent.config.workspace.base_dir.clone(), + memory_tool_search_limit: agent.config.memory.tool_search_limit, + auto_route_shell: agent.config.workspace.auto_route_shell, + is_teammate: agent.teammate_identity.is_some(), + }; + let observer = AgentObserver { + events: agent.event_bus.clone(), + snapshot_tx: agent.snapshot_tx.clone(), + snapshot: Arc::clone(&agent.snapshot), + }; + agent + .runtime + .register_agent(&agent.id, &agent.name, execution_config, &observer)?; + agent.register_tool_result_pager(); + agent.refresh_tasks_from_disk()?; + Ok(agent) + } + + /// Returns the agent's display name. + pub fn name(&self) -> &str { + &self.name + } + + /// Returns the stable persisted agent identifier. + pub fn id(&self) -> &str { + &self.id + } + + /// Returns the model identifier used by the agent. + pub fn model(&self) -> &str { + &self.model + } + + /// Updates the model and provider used for future turns, then persists the + /// new agent record so resumed sessions continue with the same setting. + pub fn set_model(&mut self, model: crate::ModelInfo) -> Result<(), RuntimeError> { + let provider = self + .runtime + .get_provider(Some(&model.provider)) + .ok_or_else(|| RuntimeError::ProviderNotFound(Some(model.provider.clone())))?; + self.model = model.id; + self.provider_id = provider.descriptor().id; + self.provider = provider; + self.persist_agent_record() + } + + /// Updates the reasoning options requested on future turns, then persists the + /// agent record so resumed sessions continue with the same setting. + /// + /// Mirrors [`set_model`](Self::set_model): a stateful override threaded into + /// every subsequent model request (the runner reads + /// `config.provider_request_options.reasoning` live). It composes with + /// `set_model` for **per-phase tiering** — e.g. run the gather rounds at a low + /// reasoning effort, then raise the effort (and switch to a stronger model) for + /// a final synthesis turn on the same agent, without re-spawning and losing the + /// gathered context. `None` clears any configured reasoning, restoring the + /// provider's default effort. + pub fn set_reasoning( + &mut self, + reasoning: Option, + ) -> Result<(), RuntimeError> { + self.config.provider_request_options.reasoning = reasoning; + self.persist_agent_record() + } + + /// Returns the effective agent configuration. + pub fn config(&self) -> &AgentConfig { + &self.config + } + + /// Returns the committed transcript history. + pub fn history(&self) -> &[Message] { + self.memory.history() + } + + /// Returns the canonical transcript items stored for this agent. + pub fn transcript(&self) -> &crate::AgentTranscript { + self.memory.transcript() + } + + /// The transcript entry the next turn will continue from. + pub fn leaf(&self) -> Option<&crate::transcript::EntryId> { + self.transcript().leaf() + } + + /// Returns to an earlier entry, so the next turn explores a new path from + /// there. + /// + /// The abandoned entries stay in the transcript, reachable through + /// [`children`](Self::children) — nothing is deleted, so the path just + /// left can be returned to the same way. Returns how many entries left + /// the active path. + pub fn branch_from( + &mut self, + entry: &crate::transcript::EntryId, + ) -> Result { + self.memory.branch_from(entry) + } + + /// The entries recorded as continuing from `entry`. More than one means + /// the conversation branched there. + pub fn children( + &self, + entry: &crate::transcript::EntryId, + ) -> Vec<&crate::transcript::TranscriptItem> { + self.transcript().children(entry) + } + + fn append_transcript_item(&mut self, item: TranscriptItem) -> Result<(), RuntimeError> { + self.memory.append_transcript_item(item) + } + + pub(crate) fn record_canonical_context( + &mut self, + content: impl Into, + ) -> Result<(), RuntimeError> { + self.append_transcript_item(TranscriptItem::canonical_context(Message::user( + ContentBlock::text(content.into()), + ))) + } + + pub(crate) fn record_delegation_request( + &mut self, + content: impl Into, + delegation: DelegationArtifact, + edge: Option, + ) -> Result<(), RuntimeError> { + self.append_transcript_item(TranscriptItem::delegation_request( + Message::user(ContentBlock::text(content.into())), + delegation, + edge, + )) + } + + pub(crate) fn record_delegation_result( + &mut self, + content: impl Into, + delegation: DelegationArtifact, + edge: Option, + ) -> Result<(), RuntimeError> { + self.append_transcript_item(TranscriptItem::delegation_result( + Message::user(ContentBlock::text(content.into())), + delegation, + edge, + )) + } + + pub(crate) fn memory_revision(&self) -> u64 { + self.memory.revision() + } + + pub(crate) fn memory_engine(&self) -> Arc { + self.runtime.memory_engine() + } + + /// Returns whether this agent is a persistent teammate rather than the lead agent. + pub fn is_teammate(&self) -> bool { + self.teammate_identity.is_some() + } + + pub(crate) fn tasks(&self) -> &[TaskItem] { + &self.tasks + } + + /// Returns the most recent committed message, if any. + pub fn last_message(&self) -> Option<&Message> { + self.memory.last_message() + } + + /// Subscribes to the agent's transient event stream. + pub fn subscribe_events(&self) -> broadcast::Receiver { + self.event_bus.subscribe() + } + + /// Watches the current agent snapshot for state updates. + pub fn watch_snapshot(&self) -> watch::Receiver { + self.snapshot_tx.subscribe() + } + + /// The tools this agent offers the model on the next round. + /// + /// A shaping typed turn (see [`Agent::run_to_output`]) narrows this to + /// exactly one tool — the terminal tool it generated — because the whole + /// point of that turn is that the model has nothing to decide but the + /// answer's shape. A working typed turn narrows nothing: its terminal tool + /// is admitted by [`can_use_tool`](Self::can_use_tool) like any other, so + /// it simply joins the ordinary roster. + pub(crate) fn tools(&self) -> Arc<[crate::tool::ProviderToolSpec]> { + let gate = self + .terminal_tool_gate + .lock() + .expect("terminal tool gate poisoned") + .clone(); + self.runtime + .tools() + .iter() + .filter(|tool| match &gate { + Some(gate) if !gate.keeps_tools => { + gate.tool_name == tool.name + && self.runtime.tool_is_visible_to_agent(&tool.name, &self.id) + } + _ => self.can_use_tool(&tool.name), + }) + .cloned() + .collect::>() + .into() + } + + pub(crate) fn can_use_tool(&self, name: &str) -> bool { + if !self.runtime.tool_is_visible_to_agent(name, &self.id) { + return false; + } + + if self + .terminal_tool_gate + .lock() + .expect("terminal tool gate poisoned") + .as_ref() + .is_some_and(|gate| gate.tool_name == name) + { + return true; + } + + if self.hidden_tools.contains(name) { + return false; + } + + if !self.config.tool_profile.allows(name) { + return false; + } + + if name == RuntimeIntrinsicTool::Idle.to_string() { + return self.teammate_identity.is_some(); + } + + // The pager's reader exists for the model only while there can be + // paged results to read. Registration is runtime-wide (the registry + // is keyed by tool name), so this per-agent gate — not registration — + // is what keeps the tool out of an unpaged agent's roster, even when + // a paging agent shares the same runtime. + if name == crate::tool::paging::READ_TOOL_RESULT_TOOL { + return self.config.tool_result_paging.is_some(); + } + + true + } + + pub(crate) fn runtime_handle(&self) -> RuntimeHandle { + self.runtime.clone() + } + + /// Registers the pager's reader when this agent enables paging. The tool + /// itself is stateless — it resolves both the retained results and the + /// page size from the calling agent's context — so one registration + /// serves every paging agent on the runtime, and re-registering is a + /// no-op. + fn register_tool_result_pager(&self) { + if self.config.tool_result_paging.is_some() { + self.runtime.register_tool(crate::tool::ReadToolResultTool); + } + } + + /// Retains the full text of a result that entered the transcript paged, + /// so `read_tool_result` can serve its later windows. In memory only, for + /// this agent's lifetime — a result the model never asks to continue + /// simply goes away with the agent. + pub(crate) fn record_paged_tool_result(&self, tool_use_id: &str, full: &str) { + self.paged_tool_results.record(tool_use_id, full); + } + + /// Returns a retained full result by `tool_use_id`. Only this agent's own + /// paged results are reachable: the store is per-agent, so one agent can + /// never read another's. + pub(crate) fn paged_tool_result(&self, tool_use_id: &str) -> Option> { + self.paged_tool_results.get(tool_use_id) + } + + pub(crate) fn max_rounds(&self) -> Option { + self.max_rounds + } + + /// What the model is told about choosing a tool on the next round. + /// + /// A shaping typed turn forces its terminal tool: it is the only tool on + /// the request, and the turn exists to produce that one call. A working + /// typed turn forces nothing — not the terminal tool, which would end the + /// turn before any work happened, and not the agent's own configured + /// choice, which would keep the turn from ever reaching the call that ends + /// it. Both would defeat the mode, so while one runs the choice is `Auto`. + pub(crate) fn tool_choice(&self) -> Option { + let gate = self + .terminal_tool_gate + .lock() + .expect("terminal tool gate poisoned") + .clone(); + if let Some(gate) = gate { + return Some(if gate.keeps_tools { + ToolChoice::Auto + } else { + ToolChoice::Tool { + name: gate.tool_name, + } + }); + } + + match self.config.tool_choice.clone() { + Some(ToolChoice::Tool { name }) if !self.can_use_tool(&name) => Some(ToolChoice::Auto), + other => other, + } + } +} diff --git a/vendor/mentra/src/agent/compact.rs b/vendor/mentra/src/agent/compact.rs new file mode 100644 index 0000000..0132ed7 --- /dev/null +++ b/vendor/mentra/src/agent/compact.rs @@ -0,0 +1,201 @@ +use crate::memory::journal::CompactionOutcome; +use crate::{ + ContentBlock, Message, + agent::AgentEvent, + compaction::compaction_request_from_agent, + error::{ErrorCategory, RuntimeError}, + memory::{ + estimated_request_tokens, micro_compact_history, required_tail_start_for_continuation, + }, +}; + +use super::{Agent, CompactionDetails, CompactionTrigger}; + +const AUTO_COMPACT_MAX_ATTEMPTS: u32 = 3; +const AUTO_COMPACT_RETRY_DELAY_MS: u64 = 500; + +impl Agent { + pub(crate) fn micro_compacted_history(&self) -> Vec { + micro_compact_history( + self.history(), + self.config.compaction.keep_recent_tool_results, + ) + } + + pub(crate) fn estimated_request_tokens(&self, messages: &[Message]) -> usize { + estimated_request_tokens(messages, self.effective_system_prompt().as_deref()) + } + + pub(crate) async fn auto_compact_if_needed(&mut self) -> Result<(), RuntimeError> { + let Some(threshold) = self.config.compaction.auto_compact_threshold_tokens else { + return Ok(()); + }; + + let messages = self.micro_compacted_history(); + if self.estimated_request_tokens(&messages) <= threshold { + return Ok(()); + } + + let preserve_from = required_tail_start_for_continuation(self.history()); + + for attempt in 1..=AUTO_COMPACT_MAX_ATTEMPTS { + match self + .compact_history(preserve_from, CompactionTrigger::Auto) + .await + { + Ok(_) => return Ok(()), + Err(err) + if err.category() == ErrorCategory::Retryable + && attempt < AUTO_COMPACT_MAX_ATTEMPTS => + { + self.emit_event(AgentEvent::RetryAttempt { + agent_id: self.id().to_string(), + error_message: err.to_string(), + attempt, + max_attempts: AUTO_COMPACT_MAX_ATTEMPTS, + next_delay_ms: AUTO_COMPACT_RETRY_DELAY_MS, + }); + tokio::time::sleep(tokio::time::Duration::from_millis( + AUTO_COMPACT_RETRY_DELAY_MS, + )) + .await; + } + Err(_) => { + // Non-retryable error or all attempts exhausted: degrade gracefully. + // The session continues with micro-compaction only. + return Ok(()); + } + } + } + + Ok(()) + } + + pub(crate) async fn compact_history( + &mut self, + preserve_from: usize, + trigger: CompactionTrigger, + ) -> Result, RuntimeError> { + if self.history().is_empty() { + return Ok(None); + } + + // A `preserve_from` of zero used to end the attempt here. Zero means + // the protected tail is the whole transcript — a single turn that is + // itself over budget — which is precisely when compaction is most + // needed. The engine has a split-turn path for it, so let it decide + // rather than silently doing nothing. + debug_assert!(preserve_from <= self.history().len()); + + let base_revision = self.memory.revision(); + // Compaction is a provider request like any other, so it goes out on + // the transport the runtime chose. Leaving it on the request's own + // value would quietly summarize over HTTP+SSE inside a run the host + // put on a websocket. + let mut provider_request_options = self.config.provider_request_options.clone(); + crate::provider::select_responses_transport( + self.provider.as_ref(), + self.runtime.responses_transport(), + &mut provider_request_options, + )?; + let Some(proposal) = self + .runtime + .compaction_engine() + .compact( + self.provider.clone(), + compaction_request_from_agent( + self.model(), + self.transcript().clone(), + &self.config.compaction, + provider_request_options, + ), + ) + .await? + else { + return Ok(None); + }; + let transcript_path = proposal.transcript_path.clone(); + let replaced_items = proposal.replaced_items; + let preserved_items = proposal.preserved_items; + let summary = proposal.summary.clone(); + self.runtime + .emit_hook(crate::runtime::RuntimeHookEvent::MemoryCompactionProposed { + agent_id: self.id().to_string(), + base_revision, + transcript_path: transcript_path.clone(), + })?; + let applied = self.memory.try_apply_compaction( + base_revision, + CompactionOutcome { + transcript_path: proposal.transcript_path, + transcript: proposal.transcript, + }, + )?; + if !applied { + let _ = + self.runtime + .emit_hook(crate::runtime::RuntimeHookEvent::MemoryCompactionSkipped { + agent_id: self.id().to_string(), + base_revision, + }); + return Ok(None); + } + self.runtime.memory_engine().store_compaction_summary( + self.id(), + self.memory.revision(), + &summary.render_for_handoff(), + )?; + self.sync_memory_snapshot(); + let _ = self + .runtime + .emit_hook(crate::runtime::RuntimeHookEvent::MemoryCompactionApplied { + agent_id: self.id().to_string(), + base_revision, + resulting_history_len: self.transcript().len(), + }); + + let details = CompactionDetails { + trigger, + mode: proposal.mode, + agent_id: self.id().to_string(), + transcript_path, + replaced_items, + preserved_items, + preserved_user_turns: proposal.preserved_user_turns, + preserved_delegation_results: proposal.preserved_delegation_results, + resulting_transcript_len: self.transcript().len(), + extracted_facts_count: proposal.diagnostics.extracted_facts_count, + summary_preview: proposal.diagnostics.summary_preview.clone(), + }; + self.emit_event(AgentEvent::ContextCompacted { + details: details.clone(), + }); + + Ok(Some(details)) + } + + pub(crate) fn inject_teammate_identity(&self, messages: &mut Vec) { + let Some(identity) = &self.teammate_identity else { + return; + }; + if messages.len() > 5 { + return; + } + + messages.insert( + 0, + Message::user(ContentBlock::Text { + text: format!( + "You are teammate '{}' with role '{}' on the team led by '{}'. Continue your assigned work and stay in character.", + self.name, identity.role, identity.lead + ), + }), + ); + messages.insert( + 1, + Message::assistant(ContentBlock::Text { + text: format!("I am {}. Continuing.", self.name), + }), + ); + } +} diff --git a/vendor/mentra/src/agent/config.rs b/vendor/mentra/src/agent/config.rs new file mode 100644 index 0000000..ee9bd00 --- /dev/null +++ b/vendor/mentra/src/agent/config.rs @@ -0,0 +1,516 @@ +use std::{ + collections::{BTreeMap, BTreeSet}, + path::PathBuf, + time::Duration, +}; + +#[cfg(test)] +use std::sync::atomic::{AtomicU64, Ordering}; + +use serde::{Deserialize, Serialize}; + +use crate::compaction::CompactionMode; +#[cfg(test)] +use crate::provider::ToolSearchMode; +use crate::provider::{ProviderRequestOptions, ToolChoice}; + +#[cfg(test)] +static NEXT_TEST_TRANSCRIPT_DIR_ID: AtomicU64 = AtomicU64::new(1); + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TaskConfig { + pub tasks_dir: PathBuf, + pub reminder_threshold: usize, +} + +impl Default for TaskConfig { + fn default() -> Self { + Self { + tasks_dir: default_tasks_dir(), + reminder_threshold: 3, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TeamAutonomyConfig { + pub enabled: bool, + pub poll_interval: Duration, + pub idle_timeout: Duration, +} + +impl Default for TeamAutonomyConfig { + fn default() -> Self { + Self { + enabled: false, + poll_interval: Duration::from_secs(5), + idle_timeout: Duration::from_secs(60), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TeamConfig { + pub team_dir: PathBuf, + pub autonomy: TeamAutonomyConfig, +} + +impl Default for TeamConfig { + fn default() -> Self { + Self { + team_dir: default_team_dir(), + autonomy: TeamAutonomyConfig::default(), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct CompactionConfig { + pub keep_recent_tool_results: usize, + pub auto_compact_threshold_tokens: Option, + pub transcript_dir: PathBuf, + pub summary_max_input_chars: usize, + pub summary_max_output_tokens: u32, + #[serde(default)] + pub mode: CompactionMode, + pub preserve_recent_user_tokens: usize, + pub preserve_recent_delegation_results: usize, + pub max_persisted_transcripts: Option, +} + +impl Default for CompactionConfig { + fn default() -> Self { + Self { + keep_recent_tool_results: 3, + auto_compact_threshold_tokens: Some(50_000), + transcript_dir: default_transcript_dir(), + summary_max_input_chars: 80_000, + summary_max_output_tokens: 2_000, + mode: CompactionMode::LocalOnly, + preserve_recent_user_tokens: 20_000, + preserve_recent_delegation_results: 8, + max_persisted_transcripts: Some(10), + } + } +} + +pub type ContextCompactionConfig = CompactionConfig; + +/// Bounds how much of an oversized tool result enters the model's view. +/// +/// A result at or below `threshold_bytes` is inserted byte-identically to a +/// run without paging. Above it, the transcript receives the first window +/// (at most `page_bytes`, cut on a line boundary) plus a trailer naming the +/// `read_tool_result` call that returns the next window; the full result is +/// retained in memory for the life of the agent so nothing is lost. +/// +/// Paging is applied *after* the runtime's own tool-result limiter +/// (`RuntimePolicy::with_max_tool_result_bytes` / +/// `with_max_tool_result_lines`), so a `threshold_bytes` above those caps +/// never triggers — the limiter clamps the result first. Enabling paging +/// therefore means raising the policy caps to whatever a tool may legitimately +/// return and leaving them as the anti-abuse backstop. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub struct ToolResultPagingConfig { + /// Results at or below this size are inserted whole. Default 64 KiB. + pub threshold_bytes: usize, + /// Maximum bytes per inserted page/window. Default 32 KiB. + pub page_bytes: usize, +} + +impl Default for ToolResultPagingConfig { + fn default() -> Self { + Self { + threshold_bytes: 64 * 1024, + page_bytes: 32 * 1024, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct WorkspaceConfig { + pub base_dir: PathBuf, + pub auto_route_shell: bool, +} + +impl Default for WorkspaceConfig { + fn default() -> Self { + let base_dir = std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")); + Self { + base_dir, + auto_route_shell: true, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct MemoryConfig { + pub auto_recall_enabled: bool, + pub auto_recall_limit: usize, + pub auto_recall_char_budget: usize, + pub tool_search_limit: usize, + pub write_tools_enabled: bool, +} + +impl Default for MemoryConfig { + fn default() -> Self { + Self { + auto_recall_enabled: true, + auto_recall_limit: 3, + auto_recall_char_budget: 2_000, + tool_search_limit: 10, + write_tools_enabled: true, + } + } +} + +#[derive(Debug, Default, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ToolProfile { + #[serde(default)] + pub allowed_tools: Option>, + #[serde(default)] + pub hidden_tools: BTreeSet, +} + +impl ToolProfile { + pub fn all() -> Self { + Self::default() + } + + pub fn only(tools: I) -> Self + where + I: IntoIterator, + S: Into, + { + Self { + allowed_tools: Some(tools.into_iter().map(Into::into).collect()), + hidden_tools: BTreeSet::new(), + } + } + + pub fn hide(tools: I) -> Self + where + I: IntoIterator, + S: Into, + { + Self { + allowed_tools: None, + hidden_tools: tools.into_iter().map(Into::into).collect(), + } + } + + pub fn allows(&self, tool_name: &str) -> bool { + if let Some(allowed_tools) = &self.allowed_tools + && !allowed_tools.contains(tool_name) + { + return false; + } + + !self.hidden_tools.contains(tool_name) + } +} + +#[cfg(not(test))] +fn default_team_dir() -> PathBuf { + crate::default_paths::workspace_default_paths().team_dir +} + +#[cfg(test)] +fn default_team_dir() -> PathBuf { + let suffix = NEXT_TEST_TRANSCRIPT_DIR_ID.fetch_add(1, Ordering::Relaxed); + std::env::temp_dir() + .join("mentra-test-team") + .join(format!("process-{}-{suffix}", std::process::id())) +} + +#[cfg(not(test))] +fn default_transcript_dir() -> PathBuf { + crate::default_paths::workspace_default_paths().transcripts_dir +} + +#[cfg(not(test))] +fn default_tasks_dir() -> PathBuf { + crate::default_paths::workspace_default_paths().tasks_dir +} + +#[cfg(test)] +fn default_tasks_dir() -> PathBuf { + let suffix = NEXT_TEST_TRANSCRIPT_DIR_ID.fetch_add(1, Ordering::Relaxed); + std::env::temp_dir() + .join("mentra-test-tasks") + .join(format!("process-{}-{suffix}", std::process::id())) +} + +#[cfg(test)] +fn default_transcript_dir() -> PathBuf { + let suffix = NEXT_TEST_TRANSCRIPT_DIR_ID.fetch_add(1, Ordering::Relaxed); + std::env::temp_dir() + .join("mentra-test-transcripts") + .join(format!("process-{}-{suffix}", std::process::id())) +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AgentConfig { + pub system: Option, + pub tool_choice: Option, + #[serde(default)] + pub tool_profile: ToolProfile, + pub temperature: Option, + pub max_output_tokens: Option, + pub metadata: BTreeMap, + #[serde(default)] + pub provider_request_options: ProviderRequestOptions, + pub team: TeamConfig, + pub task: TaskConfig, + pub workspace: WorkspaceConfig, + #[serde(default)] + pub memory: MemoryConfig, + #[serde(alias = "context_compaction")] + pub compaction: CompactionConfig, + /// `None` (the default) preserves the unpaged behaviour exactly: every + /// tool result enters the transcript as produced, and `read_tool_result` + /// is absent from the agent's tool roster. + #[serde(default)] + pub tool_result_paging: Option, +} + +impl Default for AgentConfig { + fn default() -> Self { + Self { + system: None, + tool_choice: Some(ToolChoice::default()), + tool_profile: ToolProfile::default(), + temperature: None, + max_output_tokens: Some(8192), + metadata: BTreeMap::new(), + provider_request_options: ProviderRequestOptions::default(), + team: TeamConfig::default(), + task: TaskConfig::default(), + workspace: WorkspaceConfig::default(), + memory: MemoryConfig::default(), + compaction: CompactionConfig::default(), + tool_result_paging: None, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + use crate::provider::{ReasoningEffort, ReasoningOptions}; + + fn test_path(label: &str) -> PathBuf { + std::env::temp_dir() + .join("mentra-agent-config-tests") + .join(label) + } + + #[test] + fn explicit_paths_override_defaults() { + let tasks_dir = test_path("custom-tasks"); + let team_dir = test_path("custom-team"); + let transcript_dir = test_path("custom-transcripts"); + + let config = AgentConfig { + task: TaskConfig { + tasks_dir: tasks_dir.clone(), + ..Default::default() + }, + team: TeamConfig { + team_dir: team_dir.clone(), + ..Default::default() + }, + compaction: ContextCompactionConfig { + transcript_dir: transcript_dir.clone(), + ..Default::default() + }, + ..Default::default() + }; + + assert_eq!(config.task.tasks_dir, tasks_dir); + assert_eq!(config.team.team_dir, team_dir); + assert_eq!(config.compaction.transcript_dir, transcript_dir); + } + + #[test] + fn tool_profile_defaults_to_allowing_everything() { + let profile = ToolProfile::default(); + + assert!(profile.allows("shell")); + assert!(profile.allows("files")); + } + + #[test] + fn tool_profile_only_restricts_to_allowlist() { + let profile = ToolProfile::only(["shell", "files"]); + + assert!(profile.allows("shell")); + assert!(profile.allows("files")); + assert!(!profile.allows("task")); + } + + #[test] + fn tool_profile_hide_blocks_named_tools() { + let profile = ToolProfile::hide(["shell", "background_run"]); + + assert!(!profile.allows("shell")); + assert!(!profile.allows("background_run")); + assert!(profile.allows("files")); + } + + #[test] + fn tool_profile_respects_allowlist_and_hidden_overrides() { + let profile = ToolProfile { + allowed_tools: Some(["shell", "files"].into_iter().map(str::to_string).collect()), + hidden_tools: ["shell"].into_iter().map(str::to_string).collect(), + }; + + assert!(!profile.allows("shell")); + assert!(profile.allows("files")); + assert!(!profile.allows("task")); + } + + #[test] + fn agent_config_deserializes_without_tool_profile_field() { + let config: AgentConfig = serde_json::from_value(json!({ + "system": null, + "tool_choice": serde_json::to_value(ToolChoice::Auto).expect("serialize tool choice"), + "temperature": null, + "max_output_tokens": 8192, + "metadata": {}, + "provider_request_options": {}, + "team": TeamConfig::default(), + "task": TaskConfig::default(), + "workspace": WorkspaceConfig::default(), + "memory": MemoryConfig::default(), + "context_compaction": ContextCompactionConfig::default() + })) + .expect("deserialize config without tool profile"); + + assert_eq!(config.tool_profile, ToolProfile::default()); + } + + #[test] + fn provider_request_options_default_to_disabled_tool_search() { + let options = ProviderRequestOptions::default(); + + assert_eq!(options.tool_search_mode, ToolSearchMode::Disabled); + assert_eq!(options.reasoning, None); + } + + #[test] + fn agent_config_deserializes_without_tool_search_mode() { + let config: AgentConfig = serde_json::from_value(json!({ + "system": null, + "tool_choice": serde_json::to_value(ToolChoice::Auto).expect("serialize tool choice"), + "temperature": null, + "max_output_tokens": 8192, + "metadata": {}, + "provider_request_options": { + "responses": { + "parallel_tool_calls": true + } + }, + "team": TeamConfig::default(), + "task": TaskConfig::default(), + "workspace": WorkspaceConfig::default(), + "memory": MemoryConfig::default(), + "context_compaction": ContextCompactionConfig::default() + })) + .expect("deserialize config without tool search mode"); + + assert_eq!( + config.provider_request_options.tool_search_mode, + ToolSearchMode::Disabled + ); + assert_eq!( + config + .provider_request_options + .responses + .parallel_tool_calls, + Some(true) + ); + } + + #[test] + fn tool_result_paging_is_disabled_by_default() { + assert_eq!(AgentConfig::default().tool_result_paging, None); + } + + #[test] + fn tool_result_paging_defaults_to_64_kib_threshold_and_32_kib_pages() { + let paging = ToolResultPagingConfig::default(); + + assert_eq!(paging.threshold_bytes, 64 * 1024); + assert_eq!(paging.page_bytes, 32 * 1024); + } + + #[test] + fn agent_config_deserializes_without_tool_result_paging_field() { + let config: AgentConfig = serde_json::from_value(json!({ + "system": null, + "tool_choice": serde_json::to_value(ToolChoice::Auto).expect("serialize tool choice"), + "temperature": null, + "max_output_tokens": 8192, + "metadata": {}, + "provider_request_options": {}, + "team": TeamConfig::default(), + "task": TaskConfig::default(), + "workspace": WorkspaceConfig::default(), + "memory": MemoryConfig::default(), + "context_compaction": ContextCompactionConfig::default() + })) + .expect("deserialize config persisted before paging existed"); + + assert_eq!(config.tool_result_paging, None); + } + + #[test] + fn agent_config_round_trips_tool_result_paging() { + let config = AgentConfig { + tool_result_paging: Some(ToolResultPagingConfig { + threshold_bytes: 4_096, + page_bytes: 1_024, + }), + ..Default::default() + }; + + let restored: AgentConfig = + serde_json::from_value(serde_json::to_value(&config).expect("serialize config")) + .expect("deserialize config"); + + assert_eq!(restored.tool_result_paging, config.tool_result_paging); + } + + #[test] + fn agent_config_deserializes_reasoning_options() { + let config: AgentConfig = serde_json::from_value(json!({ + "system": null, + "tool_choice": serde_json::to_value(ToolChoice::Auto).expect("serialize tool choice"), + "temperature": null, + "max_output_tokens": 8192, + "metadata": {}, + "provider_request_options": { + "reasoning": { + "effort": "high" + } + }, + "team": TeamConfig::default(), + "task": TaskConfig::default(), + "workspace": WorkspaceConfig::default(), + "memory": MemoryConfig::default(), + "context_compaction": ContextCompactionConfig::default() + })) + .expect("deserialize config with reasoning options"); + + assert_eq!( + config.provider_request_options.reasoning, + Some(ReasoningOptions { + effort: Some(ReasoningEffort::High), + summary: None, + }) + ); + } +} diff --git a/vendor/mentra/src/agent/events.rs b/vendor/mentra/src/agent/events.rs new file mode 100644 index 0000000..094e9d3 --- /dev/null +++ b/vendor/mentra/src/agent/events.rs @@ -0,0 +1,203 @@ +use std::path::PathBuf; + +use serde::{Deserialize, Serialize}; + +use crate::{ + BackgroundTaskSummary, ContentBlock, Message, TeamMemberSummary, TeamProtocolRequestSummary, + compaction::CompactionExecutionMode, runtime::TaskItem, tool::ToolCall, +}; + +#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)] +pub enum AgentStatus { + #[default] + Idle, + AwaitingModel, + Streaming, + ExecutingTool { + id: String, + name: String, + }, + Interrupted, + Finished, + Failed(String), +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct PendingToolUseSummary { + pub id: String, + pub name: String, + pub input_json: String, + pub complete: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum SpawnedAgentStatus { + Running, + Finished, + Failed(String), +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct SpawnedAgentSummary { + pub id: String, + pub name: String, + pub model: String, + pub status: SpawnedAgentStatus, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum CompactionTrigger { + Auto, + Manual, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct CompactionDetails { + pub trigger: CompactionTrigger, + pub mode: CompactionExecutionMode, + pub agent_id: String, + pub transcript_path: PathBuf, + pub replaced_items: usize, + pub preserved_items: usize, + pub preserved_user_turns: usize, + pub preserved_delegation_results: usize, + pub resulting_transcript_len: usize, + pub extracted_facts_count: usize, + pub summary_preview: String, +} + +pub type ContextCompactionTrigger = CompactionTrigger; +pub type ContextCompactionDetails = CompactionDetails; + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct AgentSnapshot { + pub status: AgentStatus, + /// Monotonic generation of the run currently reflected by this snapshot. + /// Incremented when a new `Agent::run` checkpoint has started. + #[serde(default, skip_serializing_if = "is_zero")] + pub run_generation: u64, + pub history_len: usize, + pub current_text: String, + pub pending_tool_uses: Vec, + pub pending_team_messages: usize, + pub tasks: Vec, + pub subagents: Vec, + pub teammates: Vec, + pub protocol_requests: Vec, + pub background_tasks: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum AgentEvent { + RunStarted, + ContextCompacted { + details: CompactionDetails, + }, + SubagentSpawned { + agent: SpawnedAgentSummary, + }, + SubagentFinished { + agent: SpawnedAgentSummary, + }, + TeammateSpawned { + teammate: TeamMemberSummary, + }, + TeammateUpdated { + teammate: TeamMemberSummary, + }, + TeamProtocolRequested { + request: TeamProtocolRequestSummary, + }, + TeamProtocolResolved { + request: TeamProtocolRequestSummary, + }, + TeamInboxUpdated { + unread_count: usize, + }, + BackgroundTaskStarted { + task: BackgroundTaskSummary, + }, + BackgroundTaskFinished { + task: BackgroundTaskSummary, + }, + TextDelta { + delta: String, + full_text: String, + }, + ReasoningDelta { + delta: String, + full_text: String, + }, + ToolUseUpdated { + index: usize, + id: String, + name: String, + input_json: String, + }, + ToolUseReady { + index: usize, + call: ToolCall, + }, + ToolExecutionStarted { + call: ToolCall, + }, + ToolExecutionFinished { + result: ContentBlock, + }, + AssistantMessageCommitted { + message: Message, + }, + /// Token usage from a completed model response. + UsageReport { + input_tokens: u64, + output_tokens: u64, + cache_read_tokens: u64, + cache_creation_tokens: u64, + }, + RunFinished, + ToolExecutionProgress { + id: String, + name: String, + progress: String, + }, + RetryAttempt { + agent_id: String, + error_message: String, + attempt: u32, + max_attempts: u32, + next_delay_ms: u64, + }, + RunFailed { + error: String, + }, +} + +fn is_zero(value: &u64) -> bool { + *value == 0 +} + +#[cfg(test)] +mod tests { + use serde_json::Value; + + use super::AgentSnapshot; + + #[test] + fn zero_run_generation_uses_the_pre_field_json_shape() { + let json = serde_json::to_value(AgentSnapshot::default()).expect("serialize snapshot"); + assert!(json.get("run_generation").is_none()); + + let restored: AgentSnapshot = serde_json::from_value(json).expect("load old snapshot JSON"); + assert_eq!(restored.run_generation, 0); + + let current = AgentSnapshot { + run_generation: 1, + ..AgentSnapshot::default() + }; + let Value::Object(current) = serde_json::to_value(current).expect("serialize generation") + else { + panic!("snapshot must serialize as an object"); + }; + assert_eq!(current.get("run_generation"), Some(&Value::from(1))); + } +} diff --git a/vendor/mentra/src/agent/lifecycle.rs b/vendor/mentra/src/agent/lifecycle.rs new file mode 100644 index 0000000..7170d14 --- /dev/null +++ b/vendor/mentra/src/agent/lifecycle.rs @@ -0,0 +1,136 @@ +use crate::{ContentBlock, Message, Role, error::RuntimeError, runtime::RunOptions}; + +use super::{Agent, AgentEvent, AgentStatus, TurnRunner}; + +impl Agent { + /// Sends a user turn using default run options. + pub async fn send( + &mut self, + content: impl Into>, + ) -> Result { + self.run(content, RunOptions::default()).await + } + + /// Replays the most recent failed or interrupted user turn using default run options. + pub async fn resume(&mut self) -> Result { + self.resume_with_options(RunOptions::default()).await + } + + /// Replays the most recent failed or interrupted user turn with explicit execution options. + pub async fn resume_with_options( + &mut self, + options: RunOptions, + ) -> Result { + let content = self + .memory + .resumable_user_message() + .ok_or(RuntimeError::NoResumableTurn)? + .content + .clone(); + self.run(content, options).await + } + + /// Runs a user turn with explicit execution limits and cancellation settings. + pub async fn run( + &mut self, + content: impl Into>, + options: RunOptions, + ) -> Result { + self.idle_requested = false; + self.refresh_tasks_from_disk()?; + let tasks_before_run = self.tasks.clone(); + let rounds_before_run = self.rounds_since_task; + let task_disk_state = self.capture_task_disk_state()?; + let run_id = self.start_run_checkpoint()?; + self.memory.begin_run( + run_id, + Message { + role: Role::User, + content: content.into(), + }, + )?; + self.mutate_snapshot(|snapshot| { + snapshot.run_generation = snapshot.run_generation.saturating_add(1); + snapshot.status = AgentStatus::AwaitingModel; + }); + self.sync_memory_snapshot(); + self.emit_event(AgentEvent::RunStarted); + + match TurnRunner::new(self, options).run().await { + Ok(()) => { + let final_message = self + .memory + .last_message() + .cloned() + .filter(|message| message.role == Role::Assistant); + let run_delta = self.memory.current_run_delta().unwrap_or_default(); + let finalization = (|| { + self.memory.finish_run()?; + self.sync_memory_snapshot(); + self.set_status(AgentStatus::Finished); + self.persist_agent_record()?; + self.finish_run_checkpoint() + })(); + if let Err(error) = finalization { + return Err(self.handle_finalization_error(error)); + } + + self.clear_inflight_steering(); + self.clear_inflight_team_messages(); + self.clear_inflight_background_notifications(); + self.runtime + .memory_engine() + .schedule_ingest(crate::memory::IngestRequest { + agent_id: self.id().to_string(), + source_revision: self.memory.revision(), + messages: run_delta, + }); + self.emit_event(AgentEvent::RunFinished); + final_message.ok_or(RuntimeError::EmptyAssistantResponse) + } + Err(error) => { + self.idle_requested = false; + self.requeue_inflight_steering(); + self.requeue_inflight_team_messages()?; + self.requeue_inflight_background_notifications(); + self.restore_task_state(tasks_before_run, rounds_before_run, &task_disk_state)?; + self.memory.rollback_failed_run()?; + self.sync_memory_snapshot(); + let message = error.to_string(); + self.set_status(AgentStatus::Failed(message.clone())); + self.persist_agent_record()?; + self.fail_run_checkpoint(&message)?; + self.emit_event(AgentEvent::RunFailed { error: message }); + Err(error) + } + } + } + + fn handle_finalization_error(&mut self, error: RuntimeError) -> RuntimeError { + self.idle_requested = false; + self.requeue_inflight_steering(); + let team_requeue = self.requeue_inflight_team_messages(); + self.requeue_inflight_background_notifications(); + + let message = error.to_string(); + self.set_status(AgentStatus::Failed(message.clone())); + let _ = self.persist_agent_record(); + let _ = self.fail_run_checkpoint(&message); + self.emit_event(AgentEvent::RunFailed { error: message }); + + match team_requeue { + Ok(()) => error, + Err(requeue_error) => RuntimeError::Store(format!( + "{error}; additionally failed to requeue the team inbox: {requeue_error}" + )), + } + } + + pub(crate) fn request_idle(&mut self) { + self.idle_requested = true; + } + + pub(crate) fn take_idle_requested(&mut self) -> bool { + std::mem::take(&mut self.idle_requested) + } +} diff --git a/vendor/mentra/src/agent/pending.rs b/vendor/mentra/src/agent/pending.rs new file mode 100644 index 0000000..40e477c --- /dev/null +++ b/vendor/mentra/src/agent/pending.rs @@ -0,0 +1,269 @@ +use std::collections::BTreeMap; + +use crate::{ + ContentBlock, Message, Role, + error::RuntimeError, + provider::{ContentBlockDelta, ProviderEvent, TokenUsage}, + tool::ToolCall, +}; + +use super::{AgentEvent, PendingToolUseSummary, pending_block::PendingContentBlock}; + +#[derive(Debug, Clone, Default)] +pub struct PendingAssistantTurn { + id: Option, + model: Option, + role: Option, + blocks: BTreeMap, + invalid_tool_uses: Vec, + current_text: String, + current_reasoning: String, + stop_reason: Option, + usage: Option, + stopped: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct InvalidToolUse { + pub index: usize, + pub id: String, + pub name: String, + pub input_json: String, + pub error: String, +} + +impl PendingAssistantTurn { + pub fn apply(&mut self, event: ProviderEvent) -> Result, RuntimeError> { + let mut derived_events = Vec::new(); + + match event { + ProviderEvent::ResponseHeaders(_) + | ProviderEvent::ResponseCreated + | ProviderEvent::ReasoningSummaryDelta { .. } + | ProviderEvent::ReasoningContentDelta { .. } + | ProviderEvent::ReasoningSummaryPartAdded { .. } => {} + ProviderEvent::MessageStarted { id, model, role } => { + self.id = Some(id); + self.model = Some(model); + self.role = Some(role); + } + ProviderEvent::ContentBlockStarted { index, kind } => { + self.blocks.insert(index, PendingContentBlock::from(kind)); + } + ProviderEvent::ContentBlockDelta { index, delta } => { + let block = self.blocks.get_mut(&index).ok_or_else(|| { + RuntimeError::MalformedProviderEvent(format!( + "content block delta received before start for index {index}" + )) + })?; + + match (block, delta) { + (PendingContentBlock::Text { text, .. }, ContentBlockDelta::Text(delta)) => { + text.push_str(&delta); + self.current_text.push_str(&delta); + derived_events.push(AgentEvent::TextDelta { + delta, + full_text: self.current_text.clone(), + }); + } + ( + PendingContentBlock::Thinking { thinking, .. }, + ContentBlockDelta::ThinkingText(delta), + ) => { + thinking.push_str(&delta); + self.current_reasoning.push_str(&delta); + derived_events.push(AgentEvent::ReasoningDelta { + delta, + full_text: self.current_reasoning.clone(), + }); + } + ( + PendingContentBlock::Thinking { signature, .. }, + ContentBlockDelta::ThinkingSignature(delta), + ) => { + signature.get_or_insert_with(String::new).push_str(&delta); + } + ( + PendingContentBlock::Thinking { + encrypted_content, .. + }, + ContentBlockDelta::ThinkingEncryptedContent(value), + ) => { + *encrypted_content = Some(value); + } + ( + PendingContentBlock::ToolUse { + id, + name, + input_json, + .. + }, + ContentBlockDelta::ToolUseInputJson(delta), + ) => { + input_json.push_str(&delta); + derived_events.push(AgentEvent::ToolUseUpdated { + index, + id: id.clone(), + name: name.clone(), + input_json: input_json.clone(), + }); + } + ( + PendingContentBlock::ToolResult { content, .. }, + ContentBlockDelta::ToolResultContent(delta), + ) => match (content, delta) { + ( + mentra_provider::ToolResultContent::Text(content), + mentra_provider::ToolResultContent::Text(delta), + ) => { + content.push_str(&delta); + } + (content, delta) => { + *content = delta; + } + }, + (block, delta) => { + if !block.apply_hosted_delta(&delta) { + return Err(RuntimeError::MalformedProviderEvent(format!( + "delta {delta:?} is not valid for block {}", + block.kind_name() + ))); + } + } + } + } + ProviderEvent::ContentBlockStopped { index } => { + let block = self.blocks.get_mut(&index).ok_or_else(|| { + RuntimeError::MalformedProviderEvent(format!( + "content block stop received before start for index {index}" + )) + })?; + block.mark_complete(); + + if let PendingContentBlock::ToolUse { + id, + name, + input_json, + .. + } = block + { + match serde_json::from_str(input_json) { + Ok(input) => { + derived_events.push(AgentEvent::ToolUseReady { + index, + call: ToolCall { + id: id.clone(), + name: name.clone(), + input, + }, + }); + } + Err(source) => { + self.invalid_tool_uses.push(InvalidToolUse { + index, + id: id.clone(), + name: name.clone(), + input_json: input_json.clone(), + error: source.to_string(), + }); + } + } + } + } + ProviderEvent::MessageDelta { stop_reason, usage } => { + self.stop_reason = stop_reason; + self.usage = usage; + } + ProviderEvent::MessageStopped => self.stopped = true, + } + + Ok(derived_events) + } + + pub fn to_message(&self) -> Result { + if !self.stopped { + return Err(RuntimeError::MalformedProviderEvent( + "assistant turn ended before MessageStopped".to_string(), + )); + } + + let role = self.role.clone().ok_or_else(|| { + RuntimeError::MalformedProviderEvent("assistant turn missing role".to_string()) + })?; + let mut content = Vec::with_capacity(self.blocks.len()); + + for (index, block) in &self.blocks { + if !block.is_complete() { + return Err(RuntimeError::MalformedProviderEvent(format!( + "content block {index} did not complete" + ))); + } + match block { + PendingContentBlock::ToolUse { + id, + name, + input_json, + .. + } => { + if let Ok(input) = serde_json::from_str(input_json) { + content.push(ContentBlock::ToolUse { + id: id.clone(), + name: name.clone(), + input, + }); + } + } + _ => content.push(block.to_content_block()?), + } + } + + Ok(Message { role, content }) + } + + pub fn ready_tool_calls(&self) -> Result, RuntimeError> { + let mut tool_calls = Vec::new(); + + for block in self.blocks.values() { + if let PendingContentBlock::ToolUse { + id, + name, + input_json, + complete, + } = block + && *complete + && let Ok(input) = serde_json::from_str(input_json) + { + tool_calls.push(ToolCall { + id: id.clone(), + name: name.clone(), + input, + }); + } + } + + Ok(tool_calls) + } + + pub fn pending_tool_use_summaries(&self) -> Vec { + self.blocks + .values() + .filter_map(PendingContentBlock::tool_use_summary) + .collect() + } + + pub(crate) fn invalid_tool_uses(&self) -> &[InvalidToolUse] { + &self.invalid_tool_uses + } + + pub fn current_text(&self) -> &str { + &self.current_text + } + + pub fn usage(&self) -> Option<&TokenUsage> { + self.usage.as_ref() + } + + pub fn stop_reason(&self) -> Option<&str> { + self.stop_reason.as_deref() + } +} diff --git a/vendor/mentra/src/agent/pending_block.rs b/vendor/mentra/src/agent/pending_block.rs new file mode 100644 index 0000000..01b8af9 --- /dev/null +++ b/vendor/mentra/src/agent/pending_block.rs @@ -0,0 +1,309 @@ +use crate::{ContentBlock, ImageSource, error::RuntimeError, provider::ContentBlockStart}; +use mentra_provider::{ + HostedToolSearchCall, HostedWebSearchCall, ImageGenerationCall, ImageGenerationResult, + ReasoningProvenance, ToolResultContent, WebSearchAction, +}; + +use super::PendingToolUseSummary; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(super) enum PendingContentBlock { + Text { + text: String, + complete: bool, + }, + Thinking { + thinking: String, + signature: Option, + encrypted_content: Option, + id: Option, + provenance: Option, + redacted: bool, + complete: bool, + }, + Image { + source: ImageSource, + complete: bool, + }, + ToolUse { + id: String, + name: String, + input_json: String, + complete: bool, + }, + ToolResult { + tool_use_id: String, + content: ToolResultContent, + is_error: bool, + complete: bool, + }, + HostedToolSearch { + call: HostedToolSearchCall, + complete: bool, + }, + HostedWebSearch { + call: HostedWebSearchCall, + complete: bool, + }, + ImageGeneration { + call: ImageGenerationCall, + complete: bool, + }, +} + +impl PendingContentBlock { + pub(super) fn is_complete(&self) -> bool { + match self { + PendingContentBlock::Text { complete, .. } + | PendingContentBlock::Thinking { complete, .. } + | PendingContentBlock::Image { complete, .. } + | PendingContentBlock::ToolUse { complete, .. } + | PendingContentBlock::ToolResult { complete, .. } + | PendingContentBlock::HostedToolSearch { complete, .. } + | PendingContentBlock::HostedWebSearch { complete, .. } + | PendingContentBlock::ImageGeneration { complete, .. } => *complete, + } + } + + pub(super) fn mark_complete(&mut self) { + match self { + PendingContentBlock::Text { complete, .. } + | PendingContentBlock::Thinking { complete, .. } + | PendingContentBlock::Image { complete, .. } + | PendingContentBlock::ToolUse { complete, .. } + | PendingContentBlock::ToolResult { complete, .. } + | PendingContentBlock::HostedToolSearch { complete, .. } + | PendingContentBlock::HostedWebSearch { complete, .. } + | PendingContentBlock::ImageGeneration { complete, .. } => *complete = true, + } + } + + pub(super) fn to_content_block(&self) -> Result { + match self { + PendingContentBlock::Text { text, .. } => Ok(ContentBlock::Text { text: text.clone() }), + PendingContentBlock::Thinking { + thinking, + signature, + encrypted_content, + id, + provenance, + redacted, + .. + } => Ok(ContentBlock::Thinking { + thinking: thinking.clone(), + signature: signature.clone(), + encrypted_content: encrypted_content.clone(), + id: id.clone(), + provenance: provenance.clone(), + redacted: *redacted, + }), + PendingContentBlock::Image { source, .. } => Ok(ContentBlock::Image { + source: source.clone(), + }), + PendingContentBlock::ToolUse { + id, + name, + input_json, + .. + } => Ok(ContentBlock::ToolUse { + id: id.clone(), + name: name.clone(), + input: serde_json::from_str(input_json).map_err(|source| { + RuntimeError::InvalidToolUseInput { + id: id.clone(), + name: name.clone(), + source, + } + })?, + }), + PendingContentBlock::ToolResult { + tool_use_id, + content, + is_error, + .. + } => Ok(ContentBlock::ToolResult { + tool_use_id: tool_use_id.clone(), + content: content.clone(), + is_error: *is_error, + }), + PendingContentBlock::HostedToolSearch { call, .. } => { + Ok(ContentBlock::HostedToolSearch { call: call.clone() }) + } + PendingContentBlock::HostedWebSearch { call, .. } => { + Ok(ContentBlock::HostedWebSearch { call: call.clone() }) + } + PendingContentBlock::ImageGeneration { call, .. } => { + Ok(ContentBlock::ImageGeneration { call: call.clone() }) + } + } + } + + pub(super) fn kind_name(&self) -> &'static str { + match self { + PendingContentBlock::Text { .. } => "text", + PendingContentBlock::Thinking { .. } => "thinking", + PendingContentBlock::Image { .. } => "image", + PendingContentBlock::ToolUse { .. } => "tool_use", + PendingContentBlock::ToolResult { .. } => "tool_result", + PendingContentBlock::HostedToolSearch { .. } => "hosted_tool_search", + PendingContentBlock::HostedWebSearch { .. } => "hosted_web_search", + PendingContentBlock::ImageGeneration { .. } => "image_generation", + } + } + + pub(super) fn tool_use_summary(&self) -> Option { + match self { + PendingContentBlock::ToolUse { + id, + name, + input_json, + complete, + } => Some(PendingToolUseSummary { + id: id.clone(), + name: name.clone(), + input_json: input_json.clone(), + complete: *complete, + }), + _ => None, + } + } +} + +impl From for PendingContentBlock { + fn from(value: ContentBlockStart) -> Self { + match value { + ContentBlockStart::Text => PendingContentBlock::Text { + text: String::new(), + complete: false, + }, + ContentBlockStart::Thinking { + encrypted_content, + id, + provenance, + redacted, + } => PendingContentBlock::Thinking { + thinking: String::new(), + signature: None, + encrypted_content, + id, + provenance, + redacted, + complete: false, + }, + ContentBlockStart::Image { source } => PendingContentBlock::Image { + source, + complete: false, + }, + ContentBlockStart::ToolUse { id, name } => PendingContentBlock::ToolUse { + id, + name, + input_json: String::new(), + complete: false, + }, + ContentBlockStart::ToolResult { + tool_use_id, + is_error, + content, + .. + } => PendingContentBlock::ToolResult { + tool_use_id, + content: content.unwrap_or_default(), + is_error, + complete: false, + }, + ContentBlockStart::HostedToolSearch { call } => PendingContentBlock::HostedToolSearch { + call, + complete: false, + }, + ContentBlockStart::HostedWebSearch { call } => PendingContentBlock::HostedWebSearch { + call, + complete: false, + }, + ContentBlockStart::ImageGeneration { call } => PendingContentBlock::ImageGeneration { + call, + complete: false, + }, + } + } +} + +impl PendingContentBlock { + pub(super) fn apply_hosted_delta( + &mut self, + delta: &crate::provider::ContentBlockDelta, + ) -> bool { + match (self, delta) { + ( + PendingContentBlock::HostedToolSearch { call, .. }, + crate::provider::ContentBlockDelta::HostedToolSearchQuery(query), + ) => { + call.query = Some(query.clone()); + true + } + ( + PendingContentBlock::HostedToolSearch { call, .. }, + crate::provider::ContentBlockDelta::HostedToolSearchStatus(status), + ) => { + call.status = Some(status.clone()); + true + } + ( + PendingContentBlock::HostedWebSearch { call, .. }, + crate::provider::ContentBlockDelta::HostedWebSearchAction(action), + ) => { + call.action = Some(match action { + WebSearchAction::Search { query, queries } => WebSearchAction::Search { + query: query.clone(), + queries: queries.clone(), + }, + WebSearchAction::OpenPage { url } => { + WebSearchAction::OpenPage { url: url.clone() } + } + WebSearchAction::FindInPage { url, pattern } => WebSearchAction::FindInPage { + url: url.clone(), + pattern: pattern.clone(), + }, + }); + true + } + ( + PendingContentBlock::HostedWebSearch { call, .. }, + crate::provider::ContentBlockDelta::HostedWebSearchStatus(status), + ) => { + call.status = Some(status.clone()); + true + } + ( + PendingContentBlock::ImageGeneration { call, .. }, + crate::provider::ContentBlockDelta::ImageGenerationStatus(status), + ) => { + call.status = status.clone(); + true + } + ( + PendingContentBlock::ImageGeneration { call, .. }, + crate::provider::ContentBlockDelta::ImageGenerationRevisedPrompt(prompt), + ) => { + call.revised_prompt = Some(prompt.clone()); + true + } + ( + PendingContentBlock::ImageGeneration { call, .. }, + crate::provider::ContentBlockDelta::ImageGenerationResult(result), + ) => { + call.result = Some(match result { + ImageGenerationResult::Image { source } => ImageGenerationResult::Image { + source: source.clone(), + }, + ImageGenerationResult::ArtifactRef { artifact_id } => { + ImageGenerationResult::ArtifactRef { + artifact_id: artifact_id.clone(), + } + } + }); + true + } + _ => false, + } + } +} diff --git a/vendor/mentra/src/agent/round_strategy.rs b/vendor/mentra/src/agent/round_strategy.rs new file mode 100644 index 0000000..c992848 --- /dev/null +++ b/vendor/mentra/src/agent/round_strategy.rs @@ -0,0 +1,229 @@ +//! Per-run round strategy: a host-supplied policy invoked at the two round +//! boundaries of [`Agent::run`](crate::Agent::run). +//! +//! A [`RoundStrategy`] rides on [`RunOptions`](crate::runtime::RunOptions) and so +//! belongs to exactly one run. Its state lives and dies with that run and can +//! never leak through a pooled [`Runtime`](crate::Runtime) into another run. Its +//! absence — the `None` default on `RunOptions` — reproduces mentra's built-in +//! round loop byte-for-byte. +//! +//! The runner invokes the strategy at two commit points: after a tool round's +//! results are committed to the transcript, and after a tool-free assistant +//! message is committed but before the run returns it. At each point the strategy +//! decides how the run proceeds via [`RoundDecision`]. + +use async_trait::async_trait; + +use crate::{ContentBlock, Message, ModelInfo, ReasoningOptions}; + +/// Identifies which round boundary invoked a [`RoundStrategy`]. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum RoundBoundary { + /// A tool round's results were just committed to the transcript. The run is + /// about to advance to the next round. + ToolResultsCommitted, + /// A tool-free assistant message was just committed. The run is about to + /// return it as the final message unless the strategy injects another round. + AssistantMessageCommitted, +} + +/// Provider-neutral summary of one committed tool result. +/// +/// Exposes only what a host needs to reason about a completed tool round without +/// coupling to mentra's internal tool-result representation. +#[derive(Clone, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub struct RoundToolResult { + /// The `tool_use_id` correlating this result with its originating tool call. + pub tool_use_id: String, + /// The name of the tool that produced the result. + pub tool_name: String, + /// Whether the tool reported an error. + pub is_error: bool, +} + +/// How a [`RoundAdjustment`] changes reasoning settings for subsequent rounds. +#[derive(Clone, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub enum ReasoningChange { + /// Set reasoning to the given options. + Set(ReasoningOptions), + /// Clear any configured reasoning, restoring the provider's default effort. + Clear, +} + +/// A model and/or reasoning override applied to the run's subsequent rounds. +/// +/// Applying an adjustment reuses the same live-config mechanics as +/// [`Agent::set_model`](crate::Agent::set_model) and +/// [`Agent::set_reasoning`](crate::Agent::set_reasoning): the change takes effect +/// on the next model request and persists for the remainder of the run (and, under +/// a persisting store, on the agent record) exactly as those methods do. Build one +/// with [`RoundAdjustment::new`] and the `with_*` methods. +#[derive(Clone, Debug, Default)] +#[non_exhaustive] +pub struct RoundAdjustment { + pub(crate) model: Option, + pub(crate) reasoning: Option, +} + +impl RoundAdjustment { + /// An adjustment that changes nothing. + pub fn new() -> Self { + Self::default() + } + + /// Switch the model (and its provider) for subsequent rounds. + pub fn with_model(mut self, model: ModelInfo) -> Self { + self.model = Some(model); + self + } + + /// Change the reasoning settings for subsequent rounds. + pub fn with_reasoning(mut self, reasoning: ReasoningChange) -> Self { + self.reasoning = Some(reasoning); + self + } + + /// Whether this adjustment would change nothing. + pub fn is_empty(&self) -> bool { + self.model.is_none() && self.reasoning.is_none() + } +} + +/// The decision a [`RoundStrategy`] returns at a round boundary. +/// +/// Construct one with [`RoundDecision::proceed`], [`RoundDecision::inject`], or +/// [`RoundDecision::stop`] for the common cases, or the variants directly to carry +/// a [`RoundAdjustment`]. +#[non_exhaustive] +pub enum RoundDecision { + /// Proceed with the default flow. At [`RoundBoundary::AssistantMessageCommitted`] + /// this accepts the terminal message and lets the run return; at + /// [`RoundBoundary::ToolResultsCommitted`] it advances to the next round. Any + /// carried [`RoundAdjustment`] applies to subsequent rounds. + Continue(RoundAdjustment), + /// Do not end the run: append `content` as a corrective user turn and run + /// another round. At [`RoundBoundary::AssistantMessageCommitted`] this prevents + /// the run from returning. Any carried [`RoundAdjustment`] applies to that round. + Inject { + /// Corrective content appended to the transcript as a user message. + content: Vec, + /// Model/reasoning override applied before the injected round. + adjust: RoundAdjustment, + }, + /// End the run gracefully at this boundary, committing the transcript exactly + /// as [`RunOptions::stop`](crate::runtime::RunOptions) does. A stop request + /// asserts nothing about whether the run produced a valid answer. + Stop, +} + +impl RoundDecision { + /// Continue with no model/reasoning change. + pub fn proceed() -> Self { + RoundDecision::Continue(RoundAdjustment::default()) + } + + /// Inject corrective context and run another round, with no adjustment. + pub fn inject(content: impl Into>) -> Self { + RoundDecision::Inject { + content: content.into(), + adjust: RoundAdjustment::default(), + } + } + + /// Request a graceful stop. + pub fn stop() -> Self { + RoundDecision::Stop + } +} + +/// Read-only view of a round boundary handed to a [`RoundStrategy`]. +pub struct RoundContext<'a> { + boundary: RoundBoundary, + assistant_message: Option<&'a Message>, + tool_results: &'a [RoundToolResult], + rounds_completed: usize, + model_requests: usize, + transport_retries: usize, +} + +impl<'a> RoundContext<'a> { + pub(crate) fn new( + boundary: RoundBoundary, + assistant_message: Option<&'a Message>, + tool_results: &'a [RoundToolResult], + rounds_completed: usize, + model_requests: usize, + transport_retries: usize, + ) -> Self { + Self { + boundary, + assistant_message, + tool_results, + rounds_completed, + model_requests, + transport_retries, + } + } + + /// Which boundary invoked the strategy. + pub fn boundary(&self) -> RoundBoundary { + self.boundary + } + + /// The assistant message just committed. Present only at + /// [`RoundBoundary::AssistantMessageCommitted`], and absent even there if the + /// terminal turn carried no content. + pub fn assistant_message(&self) -> Option<&Message> { + self.assistant_message + } + + /// Summaries of the tool results just committed. Present only at + /// [`RoundBoundary::ToolResultsCommitted`]; empty otherwise. + pub fn tool_results(&self) -> &[RoundToolResult] { + self.tool_results + } + + /// The number of rounds the run has entered so far, including this one. + pub fn rounds_completed(&self) -> usize { + self.rounds_completed + } + + /// The number of provider requests the run has issued so far, including + /// transient transport retries. This mirrors mentra's request counter, which + /// is not a pure logical-round counter. + pub fn model_requests(&self) -> usize { + self.model_requests + } + + /// The number of transient transport-connection retries the run has made so + /// far — the subset of [`model_requests`](Self::model_requests) that were + /// *not* the round's successful attempt. Kept distinct from + /// [`rounds_completed`](Self::rounds_completed), which counts only completed + /// logical rounds: a round that needed retries before its connection opened + /// still counts as exactly one completed round. + pub fn transport_retries(&self) -> usize { + self.transport_retries + } +} + +/// A host-supplied policy invoked at each round boundary of a single +/// [`Agent::run`](crate::Agent::run) invocation. +/// +/// The strategy is carried on [`RunOptions`](crate::runtime::RunOptions) and so is +/// bound to exactly one run: its state lives and dies with that run and can never +/// leak through a pooled [`Runtime`](crate::Runtime) into another run. Its absence +/// — the `None` default — reproduces mentra's built-in round loop exactly. +/// +/// At each boundary the strategy may [continue](RoundDecision::Continue), +/// [inject](RoundDecision::Inject) corrective context, switch the next round's +/// model or reasoning via a [`RoundAdjustment`], or request a graceful +/// [stop](RoundDecision::Stop). Every decision still passes through the run's +/// existing budget, cancellation, and deadline checks; an injected round is a +/// normal round in every respect. +#[async_trait] +pub trait RoundStrategy: Send + Sync { + /// Decide how the run should proceed at the given boundary. + async fn on_round(&self, ctx: RoundContext<'_>) -> RoundDecision; +} diff --git a/vendor/mentra/src/agent/runner.rs b/vendor/mentra/src/agent/runner.rs new file mode 100644 index 0000000..9ba8d73 --- /dev/null +++ b/vendor/mentra/src/agent/runner.rs @@ -0,0 +1,782 @@ +use std::{borrow::Cow, collections::HashMap, sync::Arc, time::Duration}; + +use crate::{ + ContentBlock, Message, Role, + background::BackgroundNotification, + error::RuntimeError, + memory::journal::PendingTurnState, + memory::{MemorySearchMode, MemorySearchRequest, build_search_query, recalled_memory_message}, + provider::Request, + runtime::{EarlyEnd, RunOptions, RuntimeHookEvent, control::is_transient_provider_error}, + team::format_inbox, + tool::{ToolCall, ToolRuntime}, + transcript::{DelegationArtifact, DelegationKind, DelegationStatus}, +}; + +use super::{ + Agent, AgentEvent, AgentStatus, PendingAssistantTurn, + pending::InvalidToolUse, + round_strategy::{ + ReasoningChange, RoundAdjustment, RoundBoundary, RoundContext, RoundDecision, + RoundStrategy, RoundToolResult, + }, +}; + +/// How the round loop should proceed after a [`RoundStrategy`] decision. +enum RoundFlow { + /// Advance to the next round (or, at the assistant boundary, return). + Continue, + /// A corrective turn was injected; run another round. + Inject, + /// End the run gracefully at this boundary. + Stop, +} + +const MEMORY_SEARCH_TIMEOUT: Duration = Duration::from_millis(250); + +pub(super) struct TurnRunner<'a> { + agent: &'a mut Agent, + options: RunOptions, + model_requests: usize, + /// Transient transport-connection retries made so far, counted separately + /// from `model_requests` (which includes them) and from the loop's `rounds` + /// counter (which counts only completed logical rounds). See + /// [`RoundContext::transport_retries`](crate::agent::RoundContext::transport_retries). + transport_retries: usize, + tool_runtime: ToolRuntime, +} + +struct StreamedTurn { + attempt: usize, + pending: PendingAssistantTurn, +} + +impl<'a> TurnRunner<'a> { + pub(super) fn new(agent: &'a mut Agent, options: RunOptions) -> Self { + let tool_runtime = ToolRuntime::new(agent); + Self { + agent, + options, + model_requests: 0, + transport_retries: 0, + tool_runtime, + } + } + + pub(super) async fn run(&mut self) -> Result<(), RuntimeError> { + let mut rounds = 0usize; + + loop { + self.options.check_limits()?; + // Graceful stop: end the run successfully at this round boundary (where + // the transcript is consistent), keeping the committed work, rather than + // failing and rolling back the way `cancellation` does. Lets a caller + // stop gathering once enough is done while preserving the context for a + // follow-up turn on the same agent. Recorded on the way out because a + // successful return says nothing about which of these two boundary + // checks produced it. + if self.options.stop_requested() { + self.options.record_early_end(EarlyEnd::StopRequested); + return Ok(()); + } + // Soft aggregate token bound: usage is only known once a round has + // completed, so this never preempts a round in progress — it only + // refuses to start another one once the last round's reported usage + // pushed the cumulative total to or past the bound. The transcript + // through the last completed round stays committed, exactly like + // `stop` above. Reached only when no stop was requested, which is + // the precedence `EarlyEnd::StopRequested` documents. + if self.options.token_budget_exceeded() { + self.options.record_early_end(EarlyEnd::TokenBudget); + return Ok(()); + } + if let Some(limit) = self.agent.max_rounds() + && rounds >= limit + { + return Err(RuntimeError::MaxRoundsExceeded(limit)); + } + + rounds += 1; + self.agent.update_run_state("awaiting_model", None)?; + let streamed = self.stream_turn().await?; + let invalid_tool_uses = streamed.pending.invalid_tool_uses().to_vec(); + if let Err(error) = self.commit_assistant_message(&streamed.pending) { + self.emit_model_response_finished( + streamed.attempt, + false, + Some(error.to_string()), + None, + None, + )?; + return Err(error); + } + let usage = streamed.pending.usage().cloned(); + self.emit_model_response_finished( + streamed.attempt, + true, + None, + streamed.pending.stop_reason().map(str::to_string), + usage.clone(), + )?; + + // Emit token usage report if available. + if let Some(ref u) = usage { + self.agent.emit_event(AgentEvent::UsageReport { + input_tokens: u.input_tokens.unwrap_or(0), + output_tokens: u.output_tokens.unwrap_or(0), + cache_read_tokens: u.cache_read_input_tokens.unwrap_or(0), + cache_creation_tokens: u.cache_creation_input_tokens.unwrap_or(0), + }); + self.options + .record_tokens(u.input_tokens.unwrap_or(0) + u.output_tokens.unwrap_or(0)); + } + + if !invalid_tool_uses.is_empty() { + self.append_invalid_tool_input_feedback(&invalid_tool_uses)?; + self.agent.note_round_without_task(); + self.agent.persist_agent_record()?; + continue; + } + + let tool_calls = streamed.pending.ready_tool_calls()?; + if tool_calls.is_empty() { + if self.agent.has_pending_steer() { + if !self.next_request_is_available(rounds)? { + self.agent.note_round_without_task(); + return Ok(()); + } + if let Some(content) = self.agent.drain_steer() { + self.inject_round_context(content)?; + self.agent.note_round_without_task(); + self.agent.persist_agent_record()?; + continue; + } + } + + if self.agent.has_pending_follow_up() { + if !self.next_request_is_available(rounds)? { + self.agent.note_round_without_task(); + return Ok(()); + } + if let Some(content) = self.agent.drain_follow_up() { + self.inject_round_context(content)?; + self.agent.note_round_without_task(); + self.agent.persist_agent_record()?; + continue; + } + } + + if let Some(strategy) = self.options.round_strategy.clone() { + let assistant_message = self + .agent + .last_message() + .filter(|message| message.role == Role::Assistant) + .cloned(); + let flow = self + .run_round_strategy( + &strategy, + RoundBoundary::AssistantMessageCommitted, + assistant_message.as_ref(), + &[], + rounds, + ) + .await?; + match flow { + RoundFlow::Inject => { + self.agent.note_round_without_task(); + self.agent.persist_agent_record()?; + continue; + } + RoundFlow::Continue | RoundFlow::Stop => { + self.agent.note_round_without_task(); + return Ok(()); + } + } + } + self.agent.note_round_without_task(); + return Ok(()); + } + + let round_strategy = self.options.round_strategy.clone(); + let call_names = round_strategy + .as_ref() + .map(|_| collect_tool_call_names(&tool_calls)); + + let execution = self + .tool_runtime + .execute_calls(self.agent, &self.options, tool_calls) + .await?; + + let tool_results = call_names + .map(|names| summarize_tool_results(&execution.results, &names)) + .unwrap_or_default(); + + let tool_result_message = Message { + role: Role::User, + content: execution.results, + }; + if execution.details.is_empty() { + self.agent.memory.append_message(tool_result_message)?; + } else { + self.agent + .memory + .append_message_with_details(tool_result_message, execution.details)?; + } + self.agent.sync_memory_snapshot(); + if execution.successful_task { + self.agent.record_task_activity(); + } else { + self.agent.note_round_without_task(); + } + self.agent.persist_agent_record()?; + if execution.end_turn { + return Ok(()); + } + + if self.agent.has_pending_steer() { + if !self.next_request_is_available(rounds)? { + return Ok(()); + } + if let Some(content) = self.agent.drain_steer() { + self.inject_round_context(content)?; + continue; + } + } + + if let Some(strategy) = round_strategy { + let flow = self + .run_round_strategy( + &strategy, + RoundBoundary::ToolResultsCommitted, + None, + &tool_results, + rounds, + ) + .await?; + if matches!(flow, RoundFlow::Stop) { + return Ok(()); + } + } + } + } + + /// Returns whether a queued injection can be followed by a provider + /// request. This check happens before draining so a graceful stop or a + /// hard run limit cannot consume queue entries that no request will see. + /// + /// The two graceful bounds are asked separately rather than as one `||` so + /// that the run ends reported as well as ended: this is a round boundary + /// like the one at the top of [`run`](Self::run), it ends the turn the same + /// way, and it therefore answers "why" the same way and in the same order. + fn next_request_is_available(&self, rounds: usize) -> Result { + self.options.check_limits()?; + if self.options.stop_requested() { + self.options.record_early_end(EarlyEnd::StopRequested); + return Ok(false); + } + if self.options.token_budget_exceeded() { + self.options.record_early_end(EarlyEnd::TokenBudget); + return Ok(false); + } + if let Some(limit) = self.agent.max_rounds() + && rounds >= limit + { + return Err(RuntimeError::MaxRoundsExceeded(limit)); + } + if self.model_requests >= self.options.model_budget() { + return Err(RuntimeError::ModelBudgetExceeded( + self.options.model_budget(), + )); + } + Ok(true) + } + + /// Invokes the per-run [`RoundStrategy`] at a round boundary and applies its + /// decision: any model/reasoning switch, then any corrective injection. The + /// returned [`RoundFlow`] tells the loop whether to continue, treat the round + /// as injected, or stop gracefully. + async fn run_round_strategy( + &mut self, + strategy: &Arc, + boundary: RoundBoundary, + assistant_message: Option<&Message>, + tool_results: &[RoundToolResult], + rounds: usize, + ) -> Result { + let context = RoundContext::new( + boundary, + assistant_message, + tool_results, + rounds, + self.model_requests, + self.transport_retries, + ); + match strategy.on_round(context).await { + RoundDecision::Continue(adjustment) => { + self.apply_round_adjustment(adjustment)?; + Ok(RoundFlow::Continue) + } + RoundDecision::Inject { content, adjust } => { + self.apply_round_adjustment(adjust)?; + self.inject_round_context(content)?; + Ok(RoundFlow::Inject) + } + RoundDecision::Stop => Ok(RoundFlow::Stop), + } + } + + /// Applies a strategy-requested model/reasoning switch to subsequent rounds, + /// reusing the live-config mechanics of [`Agent::set_model`] and + /// [`Agent::set_reasoning`] (which `stream_turn` reads live on the next round). + fn apply_round_adjustment(&mut self, adjustment: RoundAdjustment) -> Result<(), RuntimeError> { + if let Some(model) = adjustment.model { + self.agent.set_model(model)?; + } + match adjustment.reasoning { + Some(ReasoningChange::Set(options)) => self.agent.set_reasoning(Some(options))?, + Some(ReasoningChange::Clear) => self.agent.set_reasoning(None)?, + None => {} + } + Ok(()) + } + + /// Appends strategy-supplied corrective context as a committed user turn, + /// mirroring [`append_invalid_tool_input_feedback`](Self::append_invalid_tool_input_feedback) + /// so the injection is part of the replayable transcript. + fn inject_round_context(&mut self, content: Vec) -> Result<(), RuntimeError> { + self.agent.memory.append_message(Message { + role: Role::User, + content, + })?; + self.agent.sync_memory_snapshot(); + Ok(()) + } + + async fn stream_turn(&mut self) -> Result { + if self.model_requests >= self.options.model_budget() { + return Err(RuntimeError::ModelBudgetExceeded( + self.options.model_budget(), + )); + } + self.agent.inject_team_inbox()?; + self.agent.inject_background_notifications()?; + self.agent.set_status(AgentStatus::AwaitingModel); + self.agent.refresh_tasks_from_disk()?; + self.agent.auto_compact_if_needed().await?; + let provider = self.agent.provider.clone(); + let tools = self.agent.tools(); + let mut request_history = self.agent.micro_compacted_history(); + if let Some(recalled) = self.recalled_memory_message(&request_history).await { + request_history.push(recalled); + } + self.agent.inject_teammate_identity(&mut request_history); + let mut provider_request_options = self.agent.config.provider_request_options.clone(); + // Settled once, before the first attempt: a transport the provider + // cannot serve is a configuration error, and retrying it would only + // reach the same refusal five more times. + crate::provider::select_responses_transport( + provider.as_ref(), + self.agent.runtime.responses_transport(), + &mut provider_request_options, + )?; + let request = Request { + model: self.agent.model.as_str().into(), + system: self.agent.effective_system_prompt(), + messages: request_history.into(), + tools: tools.as_ref().into(), + tool_choice: self.agent.tool_choice(), + temperature: self.agent.config.temperature, + max_output_tokens: self.agent.config.max_output_tokens, + metadata: Cow::Borrowed(&self.agent.config.metadata), + provider_request_options, + }; + let mut attempt = 0usize; + let mut stream = loop { + self.options.check_limits()?; + attempt += 1; + self.model_requests += 1; + self.agent + .runtime + .emit_hook(RuntimeHookEvent::ModelRequestStarted { + agent_id: self.agent.id().to_string(), + model: self.agent.model().to_string(), + attempt, + })?; + match provider.stream(request.clone()).await { + Ok(stream) => { + self.agent + .runtime + .emit_hook(RuntimeHookEvent::ModelRequestFinished { + agent_id: self.agent.id().to_string(), + model: self.agent.model().to_string(), + attempt, + success: true, + error: None, + })?; + break stream; + } + Err(error) + if attempt <= self.options.retry_budget + && is_transient_provider_error(&error) => + { + self.transport_retries += 1; + self.agent + .runtime + .emit_hook(RuntimeHookEvent::ModelRequestFinished { + agent_id: self.agent.id().to_string(), + model: self.agent.model().to_string(), + attempt, + success: false, + error: Some(error.to_string()), + })?; + if self.model_requests >= self.options.model_budget() { + return Err(RuntimeError::ModelBudgetExceeded( + self.options.model_budget(), + )); + } + // The provider's own answer, when it gave one, beats a + // schedule that cannot know how long the window is. See + // `ProviderRetry::delay_for` for which of the two wins. + let delay = self + .options + .provider_retry + .delay_for(attempt, error.retry_after()); + self.agent.emit_event(AgentEvent::RetryAttempt { + agent_id: self.agent.id().to_string(), + error_message: error.to_string(), + attempt: attempt as u32, + max_attempts: self.options.retry_budget as u32, + next_delay_ms: delay.as_millis() as u64, + }); + tokio::time::sleep(delay).await; + continue; + } + Err(error) => { + self.agent + .runtime + .emit_hook(RuntimeHookEvent::ModelRequestFinished { + agent_id: self.agent.id().to_string(), + model: self.agent.model().to_string(), + attempt, + success: false, + error: Some(error.to_string()), + })?; + return Err(RuntimeError::FailedToStreamResponse(error)); + } + } + }; + + let mut pending = PendingAssistantTurn::default(); + self.agent.set_status(AgentStatus::Streaming); + self.agent + .memory + .update_pending_turn(Self::pending_state(&pending))?; + self.agent.sync_memory_snapshot(); + + while let Some(event) = stream.recv().await { + if let Err(error) = self.options.check_limits() { + self.emit_model_response_finished( + attempt, + false, + Some(error.to_string()), + None, + None, + )?; + return Err(error); + } + let event = match event { + Ok(event) => event, + Err(error) => { + let runtime_error = RuntimeError::FailedToStreamResponse(error); + let error_message = runtime_error.to_string(); + self.emit_model_response_finished( + attempt, + false, + Some(error_message), + None, + None, + )?; + return Err(runtime_error); + } + }; + let derived_events = match pending.apply(event) { + Ok(derived_events) => derived_events, + Err(error) => { + self.emit_model_response_finished( + attempt, + false, + Some(error.to_string()), + None, + None, + )?; + return Err(error); + } + }; + self.agent + .memory + .update_pending_turn(Self::pending_state(&pending))?; + self.agent.sync_memory_snapshot(); + + for event in derived_events { + self.agent.emit_event(event); + } + } + + Ok(StreamedTurn { attempt, pending }) + } + + fn commit_assistant_message( + &mut self, + pending: &PendingAssistantTurn, + ) -> Result<(), RuntimeError> { + let assistant_message = pending.to_message()?; + if assistant_message.content.is_empty() { + self.agent.memory.clear_pending_turn()?; + self.agent.sync_memory_snapshot(); + return Ok(()); + } + self.agent + .memory + .commit_assistant_message(assistant_message.clone())?; + self.agent.sync_memory_snapshot(); + self.agent + .emit_event(AgentEvent::AssistantMessageCommitted { + message: assistant_message, + }); + Ok(()) + } + + fn append_invalid_tool_input_feedback( + &mut self, + invalid_tool_uses: &[InvalidToolUse], + ) -> Result<(), RuntimeError> { + self.agent + .memory + .append_message(Message::user(ContentBlock::text( + format_invalid_tool_input_feedback(invalid_tool_uses), + )))?; + self.agent.sync_memory_snapshot(); + Ok(()) + } + + fn pending_state(pending: &PendingAssistantTurn) -> PendingTurnState { + PendingTurnState::new( + pending.current_text().to_string(), + pending.pending_tool_use_summaries(), + ) + } + + fn emit_model_response_finished( + &self, + attempt: usize, + success: bool, + error: Option, + stop_reason: Option, + usage: Option, + ) -> Result<(), RuntimeError> { + self.agent + .runtime + .emit_hook(RuntimeHookEvent::ModelResponseFinished { + agent_id: self.agent.id().to_string(), + model: self.agent.model().to_string(), + attempt, + success, + error, + stop_reason, + usage, + }) + } + + async fn recalled_memory_message(&self, request_history: &[Message]) -> Option { + if !self.agent.config().memory.auto_recall_enabled { + return None; + } + let query = build_search_query(request_history, self.agent.tasks()); + if query.trim().is_empty() { + return None; + } + + let memory = self.agent.runtime.memory_engine(); + let search = memory.search(MemorySearchRequest { + agent_id: self.agent.id().to_string(), + query, + limit: self.agent.config().memory.auto_recall_limit, + char_budget: Some(self.agent.config().memory.auto_recall_char_budget), + mode: MemorySearchMode::Automatic, + filter: crate::memory::MemoryListFilter::default(), + }); + let hits = match tokio::time::timeout(MEMORY_SEARCH_TIMEOUT, search).await { + Ok(Ok(hits)) => hits, + Ok(Err(_error)) => return None, + Err(_) => { + let _ = self + .agent + .runtime + .emit_hook(RuntimeHookEvent::MemorySearchFinished { + agent_id: self.agent.id().to_string(), + success: false, + result_count: 0, + error: Some("memory search timed out".to_string()), + }); + return None; + } + }; + recalled_memory_message(&hits, self.agent.config().memory.auto_recall_char_budget) + } +} + +/// Builds a `tool_use_id -> tool_name` map from the round's tool calls so a +/// committed tool result (which carries only `tool_use_id`) can be summarized +/// with its originating tool name for a [`RoundStrategy`]. +fn collect_tool_call_names(calls: &[ToolCall]) -> HashMap { + calls + .iter() + .map(|call| (call.id.clone(), call.name.clone())) + .collect() +} + +/// Summarizes committed tool-result blocks into provider-neutral +/// [`RoundToolResult`]s for a [`RoundStrategy`], correlating each result with its +/// originating tool name via `names`. +fn summarize_tool_results( + results: &[ContentBlock], + names: &HashMap, +) -> Vec { + results + .iter() + .filter_map(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + is_error, + .. + } => Some(RoundToolResult { + tool_use_id: tool_use_id.clone(), + tool_name: names.get(tool_use_id).cloned().unwrap_or_default(), + is_error: *is_error, + }), + _ => None, + }) + .collect() +} + +fn format_invalid_tool_input_feedback(invalid_tool_uses: &[InvalidToolUse]) -> String { + let mut feedback = String::from( + "One or more tool calls could not be executed because their JSON arguments were invalid. \ +Please retry with valid JSON that matches the tool schema exactly.\n\n", + ); + + for invalid in invalid_tool_uses { + feedback.push_str(&format!( + "Tool '{}' ({}) failed to parse: {}.\nRaw arguments (truncated): {}\n\n", + invalid.name, + invalid.id, + invalid.error, + truncate_tool_input(&invalid.input_json, 240) + )); + } + + feedback.truncate(feedback.trim_end().len()); + feedback +} + +fn truncate_tool_input(input: &str, max_chars: usize) -> String { + let mut truncated = input.chars().take(max_chars).collect::(); + if input.chars().count() > max_chars { + truncated.push_str("..."); + } + truncated +} + +impl Agent { + pub(super) fn inject_team_inbox(&mut self) -> Result<(), RuntimeError> { + let messages = self + .runtime + .read_team_inbox(self.config.team.team_dir.as_path(), &self.name)?; + if messages.is_empty() { + return Ok(()); + } + + self.inflight_team_messages.extend(messages.iter().cloned()); + for message in &messages { + let content = format_inbox(std::slice::from_ref(message)); + self.record_delegation_request( + content, + DelegationArtifact { + kind: DelegationKind::Teammate, + agent_id: message.sender.clone(), + agent_name: message.sender.clone(), + role: Some("teammate".to_string()), + status: DelegationStatus::Requested, + task_summary: message.content.clone(), + result_summary: None, + artifacts: Vec::new(), + }, + None, + )?; + } + self.sync_memory_snapshot(); + Ok(()) + } + + pub(super) fn clear_inflight_team_messages(&mut self) { + let _ = self + .runtime + .acknowledge_team_messages(self.config.team.team_dir.as_path(), &self.name); + self.inflight_team_messages.clear(); + } + + pub(super) fn requeue_inflight_team_messages(&mut self) -> Result<(), RuntimeError> { + let messages = std::mem::take(&mut self.inflight_team_messages); + self.runtime.requeue_team_messages( + self.config.team.team_dir.as_path(), + &self.name, + messages, + ) + } + + pub(super) fn inject_background_notifications(&mut self) -> Result<(), RuntimeError> { + let notifications = self.runtime.drain_background_notifications(&self.id); + if notifications.is_empty() { + return Ok(()); + } + + self.inflight_background_notifications + .extend(notifications.iter().cloned()); + self.record_canonical_context(format_background_results(¬ifications))?; + self.sync_memory_snapshot(); + Ok(()) + } + + pub(super) fn clear_inflight_background_notifications(&mut self) { + self.runtime.acknowledge_background_notifications(&self.id); + self.inflight_background_notifications.clear(); + } + + pub(super) fn requeue_inflight_background_notifications(&mut self) { + let notifications = std::mem::take(&mut self.inflight_background_notifications); + self.runtime + .requeue_background_notifications(&self.id, notifications); + } +} + +fn format_background_results(notifications: &[BackgroundNotification]) -> String { + let lines = notifications + .iter() + .map(|notification| { + format!( + "[bg:{}] status={} command=\"{}\" output=\"{}\"", + notification.task_id, + notification.status, + escape_background_field(¬ification.command), + escape_background_field(¬ification.output_preview), + ) + }) + .collect::>() + .join("\n"); + + format!("\n{lines}\n") +} + +fn escape_background_field(value: &str) -> String { + value.replace('\\', "\\\\").replace('"', "\\\"") +} diff --git a/vendor/mentra/src/agent/snapshot.rs b/vendor/mentra/src/agent/snapshot.rs new file mode 100644 index 0000000..629d529 --- /dev/null +++ b/vendor/mentra/src/agent/snapshot.rs @@ -0,0 +1,114 @@ +use crate::{agent::AgentEvent, runtime::PersistedAgentRecord}; + +use super::{Agent, AgentEventTapGuard, AgentStatus}; + +impl Agent { + pub(crate) fn emit_event(&self, event: AgentEvent) { + self.event_bus.send(event); + } + + pub(crate) fn event_sender(&self) -> super::AgentEventBus { + self.event_bus.clone() + } + + pub(crate) fn register_event_tap( + &self, + tap: impl Fn(&AgentEvent) + Send + Sync + 'static, + ) -> AgentEventTapGuard { + self.event_bus.register_tap(tap) + } + + pub(crate) fn set_status(&mut self, status: AgentStatus) { + self.mutate_snapshot(|snapshot| { + snapshot.status = status; + }); + } + + pub(super) fn publish_snapshot(&self) { + let snapshot = self + .snapshot + .lock() + .expect("agent snapshot poisoned") + .clone(); + self.snapshot_tx.send_replace(snapshot); + } + + pub(super) fn mutate_snapshot(&self, update: impl FnOnce(&mut crate::agent::AgentSnapshot)) { + { + let mut snapshot = self.snapshot.lock().expect("agent snapshot poisoned"); + update(&mut snapshot); + } + self.publish_snapshot(); + } + + pub(crate) fn sync_memory_snapshot(&self) { + let memory_view = self.memory.snapshot_view(); + self.mutate_snapshot(|snapshot| { + snapshot.history_len = memory_view.history_len; + snapshot.current_text = memory_view.current_text; + snapshot.pending_tool_uses = memory_view.pending_tool_uses; + }); + } + + pub(crate) fn persisted_record(&self) -> PersistedAgentRecord { + PersistedAgentRecord { + id: self.id.clone(), + runtime_identifier: self.runtime.persisted_runtime_identifier().to_string(), + name: self.name.clone(), + model: self.model.clone(), + provider_id: self.provider_id.clone(), + config: self.config.clone(), + hidden_tools: self.hidden_tools.clone(), + max_rounds: self.max_rounds, + teammate_identity: self.teammate_identity.clone(), + rounds_since_task: self.rounds_since_task, + idle_requested: self.idle_requested, + status: self.watch_snapshot().borrow().status.clone(), + subagents: self.watch_snapshot().borrow().subagents.clone(), + } + } + + pub(crate) fn persist_agent_record(&self) -> Result<(), crate::error::RuntimeError> { + self.runtime + .store() + .save_agent_record(&self.persisted_record()) + } + + pub(crate) fn start_run_checkpoint(&mut self) -> Result { + let run_id = self.runtime.store().start_run(&self.id)?; + self.current_run_id = Some(run_id.clone()); + Ok(run_id) + } + + pub(crate) fn update_run_state( + &self, + state: &str, + error: Option<&str>, + ) -> Result<(), crate::runtime::RuntimeError> { + if let Some(run_id) = &self.current_run_id { + self.runtime + .store() + .update_run_state(run_id, state, error)?; + } + Ok(()) + } + + pub(crate) fn finish_run_checkpoint(&mut self) -> Result<(), crate::runtime::RuntimeError> { + if let Some(run_id) = self.current_run_id.as_deref() { + self.runtime.store().finish_run(run_id)?; + self.current_run_id = None; + } + Ok(()) + } + + pub(crate) fn fail_run_checkpoint( + &mut self, + error: &str, + ) -> Result<(), crate::runtime::RuntimeError> { + if let Some(run_id) = self.current_run_id.as_deref() { + self.runtime.store().fail_run(run_id, error)?; + self.current_run_id = None; + } + Ok(()) + } +} diff --git a/vendor/mentra/src/agent/steering.rs b/vendor/mentra/src/agent/steering.rs new file mode 100644 index 0000000..1c39c7b --- /dev/null +++ b/vendor/mentra/src/agent/steering.rs @@ -0,0 +1,217 @@ +use std::{ + collections::VecDeque, + sync::{Arc, Mutex, MutexGuard}, +}; + +use crate::{ContentBlock, Message, error::RuntimeError, runtime::RunOptions}; + +use super::Agent; + +/// Controls how many queued entries are injected at one eligible boundary. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub enum QueueMode { + /// Drain every currently queued entry into one model round. + All, + /// Drain exactly one entry per eligible boundary. + #[default] + OneAtATime, +} + +#[derive(Default)] +struct SteeringQueues { + steer: VecDeque>, + follow_up: VecDeque>, + steer_mode: QueueMode, + follow_up_mode: QueueMode, +} + +/// Cloneable, agent-scoped handle for live steering and deferred follow-ups. +/// +/// Obtain this handle before calling [`Agent::run`](crate::Agent::run). `steer` +/// entries are eligible at either committed round boundary; `follow_up` +/// entries are eligible only when a tool-free assistant response would +/// otherwise stop the run. Queues are in-memory and survive sequential runs of +/// this agent, but are never shared with another agent on the same runtime. +#[derive(Clone, Default)] +pub struct SteeringHandle { + queues: Arc>, +} + +impl SteeringHandle { + pub(crate) fn new() -> Self { + Self::default() + } + + /// Enqueues context for the next eligible round boundary. + pub fn steer(&self, content: impl Into>) { + let content = content.into(); + if !content.is_empty() { + self.lock().steer.push_back(content); + } + } + + /// Enqueues context used only when a run would otherwise stop. + pub fn follow_up(&self, content: impl Into>) { + let content = content.into(); + if !content.is_empty() { + self.lock().follow_up.push_back(content); + } + } + + /// Removes steering entries that have not yet been injected. + pub fn clear_steer(&self) { + self.lock().steer.clear(); + } + + /// Removes follow-up entries that have not yet been injected. + pub fn clear_follow_up(&self) { + self.lock().follow_up.clear(); + } + + /// Returns whether either queue contains an entry awaiting injection. + pub fn has_pending(&self) -> bool { + let queues = self.lock(); + !queues.steer.is_empty() || !queues.follow_up.is_empty() + } + + /// Sets the steering drain mode for subsequent boundaries. + pub fn set_steer_mode(&self, mode: QueueMode) { + self.lock().steer_mode = mode; + } + + /// Sets the follow-up drain mode for subsequent would-stop boundaries. + pub fn set_follow_up_mode(&self, mode: QueueMode) { + self.lock().follow_up_mode = mode; + } + + pub(crate) fn has_steer(&self) -> bool { + !self.lock().steer.is_empty() + } + + pub(crate) fn has_follow_up(&self) -> bool { + !self.lock().follow_up.is_empty() + } + + fn drain_steer(&self) -> Vec> { + let mut queues = self.lock(); + let mode = queues.steer_mode; + drain(&mut queues.steer, mode) + } + + fn drain_follow_up(&self) -> Vec> { + let mut queues = self.lock(); + let mode = queues.follow_up_mode; + drain(&mut queues.follow_up, mode) + } + + fn prepend_steer(&self, entries: Vec>) { + prepend(&mut self.lock().steer, entries); + } + + fn prepend_follow_up(&self, entries: Vec>) { + prepend(&mut self.lock().follow_up, entries); + } + + fn lock(&self) -> MutexGuard<'_, SteeringQueues> { + self.queues + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + } +} + +impl Agent { + /// Returns an agent-scoped handle suitable for use while `run(&mut self)` + /// holds the mutable agent borrow. + pub fn steering_handle(&self) -> SteeringHandle { + self.steering.clone() + } + + /// Idle convenience for enqueueing a steer on this agent. + pub fn steer(&self, content: impl Into>) { + self.steering.steer(content); + } + + /// Idle convenience for enqueueing a would-stop follow-up on this agent. + pub fn follow_up(&self, content: impl Into>) { + self.steering.follow_up(content); + } + + /// Starts an idle run from the next queued steer. + /// + /// This is the only automatic consumption point for a steer while no run is + /// active. Follow-ups remain reserved for a running turn's would-stop + /// boundary. A failed run prepends the consumed entry back onto the queue. + pub async fn run_queued(&mut self, options: RunOptions) -> Result { + let Some(content) = self.drain_steer() else { + return Err(RuntimeError::OperationDenied( + "no queued steering input is available".to_string(), + )); + }; + + let result = self.run(content, options).await; + if result.is_err() { + // `run` normally requeues on its rollback path. This second call is + // intentionally idempotent and covers errors raised before the run + // checkpoint is established. + self.requeue_inflight_steering(); + } + result + } + + pub(super) fn has_pending_steer(&self) -> bool { + self.steering.has_steer() + } + + pub(super) fn has_pending_follow_up(&self) -> bool { + self.steering.has_follow_up() + } + + pub(super) fn drain_steer(&mut self) -> Option> { + let entries = self.steering.drain_steer(); + if entries.is_empty() { + return None; + } + let content = flatten(&entries); + self.inflight_steer.extend(entries); + Some(content) + } + + pub(super) fn drain_follow_up(&mut self) -> Option> { + let entries = self.steering.drain_follow_up(); + if entries.is_empty() { + return None; + } + let content = flatten(&entries); + self.inflight_follow_up.extend(entries); + Some(content) + } + + pub(super) fn clear_inflight_steering(&mut self) { + self.inflight_steer.clear(); + self.inflight_follow_up.clear(); + } + + pub(super) fn requeue_inflight_steering(&mut self) { + self.steering + .prepend_steer(std::mem::take(&mut self.inflight_steer)); + self.steering + .prepend_follow_up(std::mem::take(&mut self.inflight_follow_up)); + } +} + +fn drain(queue: &mut VecDeque>, mode: QueueMode) -> Vec> { + match mode { + QueueMode::All => queue.drain(..).collect(), + QueueMode::OneAtATime => queue.pop_front().into_iter().collect(), + } +} + +fn prepend(queue: &mut VecDeque>, entries: Vec>) { + for entry in entries.into_iter().rev() { + queue.push_front(entry); + } +} + +fn flatten(entries: &[Vec]) -> Vec { + entries.iter().flatten().cloned().collect() +} diff --git a/vendor/mentra/src/agent/subagent.rs b/vendor/mentra/src/agent/subagent.rs new file mode 100644 index 0000000..4a9739b --- /dev/null +++ b/vendor/mentra/src/agent/subagent.rs @@ -0,0 +1,128 @@ +use std::{borrow::Cow, collections::HashSet, sync::Arc}; + +use crate::{ + Role, + error::RuntimeError, + provider::Provider, + runtime::{RuntimeIntrinsicTool, handle::RuntimeHandle}, +}; + +use super::{ + Agent, AgentConfig, AgentSpawnOptions, SpawnedAgentStatus, SpawnedAgentSummary, + TeammateIdentity, +}; + +const SUBAGENT_MAX_ROUNDS: usize = 30; +const SUBAGENT_SYSTEM_PROMPT: &str = "You are a subagent working for another agent. Solve the delegated task, use tools when helpful, and finish with a concise final answer for the parent agent."; + +#[derive(Clone)] +pub(crate) struct DisposableSubagentTemplate { + runtime: RuntimeHandle, + model: String, + parent_name: String, + config: AgentConfig, + provider: Arc, + hidden_tools: HashSet, + teammate_identity: Option, +} + +impl DisposableSubagentTemplate { + pub(crate) fn from_agent(agent: &Agent) -> Self { + Self { + runtime: agent.runtime.clone(), + model: agent.model.clone(), + parent_name: agent.name.clone(), + config: agent.config.clone(), + provider: Arc::clone(&agent.provider), + hidden_tools: agent.hidden_tools.clone(), + teammate_identity: agent.teammate_identity.clone(), + } + } + + pub(crate) fn spawn(&self) -> Result { + let mut hidden_tools = self.hidden_tools.clone(); + hidden_tools.insert(RuntimeIntrinsicTool::Task.to_string()); + + let mut config = self.config.clone(); + config.system = Some(build_subagent_system_prompt( + self.config.system.as_deref().map(Cow::Borrowed), + )); + + Agent::new( + self.runtime.clone(), + self.model.clone(), + format!("{}::task", self.parent_name), + config, + Arc::clone(&self.provider), + AgentSpawnOptions { + hidden_tools, + max_rounds: Some(SUBAGENT_MAX_ROUNDS), + teammate_identity: self.teammate_identity.clone(), + }, + ) + } +} + +impl Agent { + pub(crate) fn spawn_subagent(&self) -> Result { + self.disposable_subagent_template().spawn() + } + + pub(crate) fn disposable_subagent_template(&self) -> DisposableSubagentTemplate { + DisposableSubagentTemplate::from_agent(self) + } + + pub(crate) fn register_subagent(&mut self, agent: &Agent) -> SpawnedAgentSummary { + let summary = SpawnedAgentSummary { + id: agent.id.clone(), + name: agent.name.clone(), + model: agent.model.clone(), + status: SpawnedAgentStatus::Running, + }; + let summary_for_snapshot = summary.clone(); + self.mutate_snapshot(|snapshot| { + snapshot.subagents.push(summary_for_snapshot); + }); + summary + } + + pub(crate) fn finish_subagent( + &mut self, + id: &str, + status: SpawnedAgentStatus, + ) -> Option { + let mut finished = None; + self.mutate_snapshot(|snapshot| { + if let Some(summary) = snapshot.subagents.iter_mut().find(|agent| agent.id == id) { + summary.status = status; + finished = Some(summary.clone()); + } + }); + finished + } + + pub(crate) fn final_text_summary(&self) -> String { + let Some(message) = self.last_message() else { + return "(no summary)".to_string(); + }; + + if message.role != Role::Assistant { + return "(no summary)".to_string(); + } + + let text = message.text(); + + if text.is_empty() { + "(no summary)".to_string() + } else { + text + } + } +} + +pub(super) fn build_subagent_system_prompt(base: Option>) -> String { + match base { + Some(system) => format!("{system}\n\n{SUBAGENT_SYSTEM_PROMPT}"), + None => SUBAGENT_SYSTEM_PROMPT.to_string(), + } +} diff --git a/vendor/mentra/src/agent/task_state.rs b/vendor/mentra/src/agent/task_state.rs new file mode 100644 index 0000000..fd9292d --- /dev/null +++ b/vendor/mentra/src/agent/task_state.rs @@ -0,0 +1,143 @@ +use std::borrow::Cow; + +use crate::error::RuntimeError; +use crate::runtime::{ + TaskBoard, TaskStateSnapshot, + task::{TASK_REMINDER_TEXT, TaskAccess, TaskIntrinsicTool, has_unfinished_tasks}, +}; + +use super::Agent; + +impl Agent { + /// Returns a task-board view with this agent's own access identity. + /// + /// Lead agents retain lead privileges. Teammates remain constrained to + /// their own tasks and cannot edit dependency edges. Task-board mutations + /// update the shared store immediately; this agent's cached snapshot is + /// refreshed at the next normal task refresh/run boundary. + pub fn task_board(&self) -> TaskBoard { + TaskBoard::agent( + self.runtime.clone(), + self.config.task.tasks_dir.clone(), + self.name.clone(), + self.teammate_identity.is_some(), + ) + } + + pub(crate) fn effective_system_prompt(&self) -> Option> { + let mut sections = Vec::new(); + + if self.rounds_since_task >= self.config.task.reminder_threshold + && has_unfinished_tasks(&self.tasks) + { + sections.push(TASK_REMINDER_TEXT.to_string()); + } + + if let Some(system) = &self.config.system { + sections.push(system.clone()); + } + + if let Some(skills) = self.runtime.skill_descriptions() { + sections.push(skills); + } + + if sections.is_empty() { + None + } else { + Some(Cow::Owned(sections.join("\n\n"))) + } + } + + pub(crate) fn note_round_without_task(&mut self) { + if has_unfinished_tasks(&self.tasks) { + self.rounds_since_task += 1; + } + } + + pub(crate) fn record_task_activity(&mut self) { + self.rounds_since_task = 0; + } + + pub(crate) fn refresh_tasks_from_disk(&mut self) -> Result<(), RuntimeError> { + let tasks = self + .runtime + .store() + .load_tasks(self.config.task.tasks_dir.as_path())?; + self.tasks = tasks; + let tasks = self.tasks.clone(); + self.mutate_snapshot(|snapshot| { + snapshot.tasks = tasks; + }); + Ok(()) + } + + pub(crate) fn task_access(&self) -> TaskAccess<'_> { + match &self.teammate_identity { + Some(_) => TaskAccess::Teammate(self.name.as_str()), + None => TaskAccess::Lead, + } + } + + pub(crate) fn try_claim_ready_task( + &mut self, + ) -> Result, RuntimeError> { + self.refresh_tasks_from_disk()?; + if self.owns_unfinished_tasks() { + return Ok(None); + } + + match self.execute_task_mutation(&TaskIntrinsicTool::Claim, serde_json::json!({})) { + Ok(content) => { + self.refresh_tasks_from_disk()?; + serde_json::from_str::(&content) + .map(Some) + .map_err(RuntimeError::FailedToSerializeTasks) + } + Err(error) if error == "No ready unowned tasks are available to claim" => Ok(None), + Err(error) => Err(RuntimeError::InvalidTask(error)), + } + } + + pub(crate) fn execute_task_mutation( + &self, + tool: &TaskIntrinsicTool, + input: serde_json::Value, + ) -> Result { + self.runtime.execute_task_mutation( + tool, + input, + self.config.task.tasks_dir.as_path(), + self.task_access(), + ) + } + + pub(super) fn capture_task_disk_state(&self) -> Result { + self.runtime + .store() + .capture_tasks(self.config.task.tasks_dir.as_path()) + } + + fn owns_unfinished_tasks(&self) -> bool { + self.tasks.iter().any(|task| { + task.owner == self.name && !matches!(task.status, crate::runtime::TaskStatus::Completed) + }) + } + + pub(super) fn restore_task_state( + &mut self, + tasks: Vec, + rounds_since_task: usize, + disk_state: &TaskStateSnapshot, + ) -> Result<(), RuntimeError> { + self.runtime + .store() + .restore_tasks(self.config.task.tasks_dir.as_path(), disk_state)?; + self.tasks = tasks; + self.rounds_since_task = rounds_since_task; + let tasks = self.tasks.clone(); + self.mutate_snapshot(|snapshot| { + snapshot.tasks = tasks; + }); + Ok(()) + } +} diff --git a/vendor/mentra/src/agent/team.rs b/vendor/mentra/src/agent/team.rs new file mode 100644 index 0000000..f3b1c6a --- /dev/null +++ b/vendor/mentra/src/agent/team.rs @@ -0,0 +1,197 @@ +use std::{borrow::Cow, sync::Arc}; + +use serde::Deserialize; +use serde_json::Value; +use tokio::sync::Mutex as AsyncMutex; + +use crate::error::RuntimeError; +use crate::runtime::task::TaskIntrinsicTool; +use crate::team::{ + TEAMMATE_MAX_ROUNDS, TeamDispatch, TeamIntrinsicTool, TeamMemberStatus, TeamMemberSummary, + TeamMessage, TeamProtocolRequestSummary, build_teammate_system_prompt, +}; + +use super::{Agent, AgentSpawnOptions, TeammateIdentity}; + +impl Agent { + pub async fn spawn_teammate( + &mut self, + name: impl Into, + role: impl Into, + prompt: Option, + ) -> Result { + let name = name.into(); + let role = role.into(); + if name.trim().is_empty() { + return Err(RuntimeError::InvalidTeam( + "Teammate name must not be empty".to_string(), + )); + } + if role.trim().is_empty() { + return Err(RuntimeError::InvalidTeam( + "Teammate role must not be empty".to_string(), + )); + } + if name == self.name { + return Err(RuntimeError::InvalidTeam( + "Teammate name must differ from the current agent".to_string(), + )); + } + + let mut hidden_tools = self.hidden_tools.clone(); + hidden_tools.extend(teammate_hidden_tools()); + + let mut config = self.config.clone(); + config.system = Some(build_teammate_system_prompt( + self.config.system.as_deref().map(Cow::Borrowed), + &name, + &role, + &self.name, + )); + + let teammate = Self::new( + self.runtime.clone(), + self.model.clone(), + name.clone(), + config, + Arc::clone(&self.provider), + AgentSpawnOptions { + hidden_tools, + max_rounds: Some(TEAMMATE_MAX_ROUNDS), + teammate_identity: Some(TeammateIdentity { + role: role.clone(), + lead: self.name.clone(), + }), + }, + )?; + + let summary = TeamMemberSummary { + id: teammate.id().to_string(), + name: name.clone(), + role, + model: teammate.model().to_string(), + status: TeamMemberStatus::Idle, + }; + + let team_dir = self.config.team.team_dir.clone(); + let actor = Arc::new(AsyncMutex::new(teammate)); + let actor_handle = self.runtime.spawn_teammate_actor(&team_dir, &name, actor)?; + + let summary = self + .runtime + .register_teammate(&team_dir, summary, actor_handle)?; + + if let Some(prompt) = prompt.filter(|prompt| !prompt.trim().is_empty()) { + self.send_team_message(&name, prompt)?; + } + + Ok(summary) + } + + pub(crate) fn revive_teammate_actor(self) -> Result<(), RuntimeError> { + let Some(identity) = self.teammate_identity.clone() else { + return Err(RuntimeError::InvalidTeam( + "Only teammate agents can be revived as teammate actors".to_string(), + )); + }; + + let summary = TeamMemberSummary { + id: self.id().to_string(), + name: self.name.clone(), + role: identity.role, + model: self.model.clone(), + status: TeamMemberStatus::Idle, + }; + + let runtime = self.runtime.clone(); + let team_dir = self.config.team.team_dir.clone(); + let actor = Arc::new(AsyncMutex::new(self)); + let actor_handle = runtime.spawn_teammate_actor(&team_dir, &summary.name, actor)?; + runtime.register_teammate(&team_dir, summary.clone(), actor_handle)?; + runtime.wake_teammate(&team_dir, &summary.name)?; + Ok(()) + } + + pub fn send_team_message( + &self, + to: &str, + content: impl Into, + ) -> Result { + self.runtime.send_team_message( + self.config.team.team_dir.as_path(), + &self.name, + to, + content.into(), + ) + } + + pub fn broadcast_team_message( + &self, + content: impl Into, + ) -> Result, RuntimeError> { + self.runtime.broadcast_team_message( + self.config.team.team_dir.as_path(), + &self.name, + content.into(), + ) + } + + pub fn read_team_inbox(&self) -> Result, RuntimeError> { + self.runtime + .read_team_inbox(self.config.team.team_dir.as_path(), &self.name) + } + + pub fn request_team_protocol( + &self, + to: &str, + protocol: impl Into, + content: impl Into, + ) -> Result { + self.runtime.create_team_request( + self.config.team.team_dir.as_path(), + &self.name, + to, + protocol.into(), + content.into(), + ) + } + + pub fn respond_team_protocol( + &self, + request_id: &str, + approve: bool, + reason: Option, + ) -> Result { + self.runtime.resolve_team_request( + self.config.team.team_dir.as_path(), + &self.name, + request_id, + approve, + reason, + ) + } +} + +#[derive(Debug, Deserialize)] +struct TaskInput { + prompt: String, +} + +fn teammate_hidden_tools() -> [String; 3] { + [ + TeamIntrinsicTool::Spawn.to_string(), + TeamIntrinsicTool::Broadcast.to_string(), + TaskIntrinsicTool::Create.to_string(), + ] +} + +pub(crate) fn parse_task_input(input: Value) -> Result { + let parsed = serde_json::from_value::(input) + .map_err(|error| format!("Invalid task input: {error}"))?; + + if parsed.prompt.trim().is_empty() { + return Err("Task prompt must not be empty".to_string()); + } + + Ok(parsed.prompt) +} diff --git a/vendor/mentra/src/agent/terminal_output.rs b/vendor/mentra/src/agent/terminal_output.rs new file mode 100644 index 0000000..ffe2154 --- /dev/null +++ b/vendor/mentra/src/agent/terminal_output.rs @@ -0,0 +1,319 @@ +use std::{ + collections::HashSet, + sync::{ + Arc, Mutex, + atomic::{AtomicU64, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use serde::de::DeserializeOwned; +use serde_json::Value; + +use crate::{ + ContentBlock, Message, Role, + error::RuntimeError, + runtime::{RunOptions, RuntimeHandle}, + tool::{ + ToolContext, ToolDefinition, ToolDurability, ToolExecutor, ToolOutput, ToolSideEffectLevel, + ToolSpec, + }, +}; + +use super::Agent; + +static NEXT_TERMINAL_TOOL_ID: AtomicU64 = AtomicU64::new(1); + +/// Provider-facing definition of a typed terminal tool. +#[derive(Debug, Clone)] +pub struct TerminalOutputSpec { + pub tool_name: String, + pub description: String, + pub schema: Value, + /// Whether the run keeps its ordinary tools while it answers. + /// + /// `false` — what [`new`](Self::new) gives you — is a *shaping* turn: the + /// generated terminal tool is the only tool the run holds, so it can only + /// put a shape on what the conversation already contains. `true` — see + /// [`with_tools`](Self::with_tools) — is a *working* turn: the run keeps + /// the agent's whole toolset and ends by calling the terminal tool. + /// [`Agent::run_to_output`] describes what each costs. + pub keeps_tools: bool, +} + +impl TerminalOutputSpec { + pub fn new( + tool_name: impl Into, + description: impl Into, + schema: Value, + ) -> Self { + Self { + tool_name: tool_name.into(), + description: description.into(), + schema, + keeps_tools: false, + } + } + + /// Lets the run work before it answers, instead of only shaping what it + /// already has. + /// + /// A shaping turn cannot read a file, run a command, or reach an MCP + /// server, so asking one for anything it has not already been told + /// produces a well-formed answer from a model that looked at nothing — + /// and reports it as a success. The way out has been to spend two turns + /// on every read-then-answer workflow: one to gather, one to shape. This + /// spends one. The run holds its ordinary tools alongside the terminal + /// tool, works as many rounds as it needs, and ends the turn by calling + /// the terminal tool with the answer. + /// + /// The cost is that nothing forces the ending: see + /// [`Agent::run_to_output`] for what a run that never calls the tool + /// returns instead. + pub fn with_tools(mut self) -> Self { + self.keeps_tools = true; + self + } +} + +/// What an in-flight [`Agent::run_to_output`] tells the rest of the agent +/// about the turn it is running: which generated tool ends it, and whether +/// the ordinary toolset is on the request beside that tool. +/// +/// Read on every round by [`Agent::tools`] and [`Agent::tool_choice`], which +/// is why it holds the mode rather than the name alone — the two answers have +/// to agree about which turn this is, and a name cannot say. +#[derive(Debug, Clone)] +pub(super) struct TerminalToolGate { + pub(super) tool_name: String, + pub(super) keeps_tools: bool, +} + +/// Typed value and committed tool-result message produced by [`Agent::run_to_output`]. +#[derive(Debug, Clone)] +pub struct FinalOutput { + pub value: T, + pub message: Message, +} + +impl Agent { + /// Runs until a generated, agent-scoped terminal tool returns a typed value. + /// + /// The helper does not use provider-level `response_format`. It registers + /// one tool whose input schema *is* the requested shape, preserves the + /// tool input as transcript `details`, and extracts it by the exact + /// `tool_use_id` from the newly committed final transcript item. + /// + /// What the run may do on its way to that call is + /// [`TerminalOutputSpec::keeps_tools`]: + /// + /// - **Shaping**, the default. The terminal tool is the only tool on the + /// request and the provider is told to call it. The turn cannot read a + /// file, run a command, or reach an MCP server, so the only thing left + /// to decide is the shape of what the conversation already holds, and + /// one round decides it. + /// - **Working**, [`TerminalOutputSpec::with_tools`]. The agent's ordinary + /// toolset is on the request beside the terminal tool and no choice is + /// forced — forcing one would preclude the very rounds that are the + /// point. The run gathers for as many rounds as it needs and ends the + /// turn by calling the terminal tool. + /// + /// Either way the terminal call ends the round it appears in: calls + /// scheduled after it in that same round are never executed, and each is + /// given an explicit `is_error` result saying so. Where the model emits + /// two terminal calls in one round, the first is the answer and the second + /// is one of those skipped calls. + /// + /// Only that call produces a value. A working run that ends any other way + /// — on prose, or at the round boundary where [`RunOptions::stop`] or + /// [`RunOptions::token_budget`] refuses another round — has nothing to + /// return and fails with `MalformedProviderEvent("run completed without + /// invoking the expected terminal tool")`, while keeping everything it + /// gathered in the transcript. [`RunOptions::ended_early`] says which + /// bound, when one was the reason. + /// + /// A run that ends on a terminal call ends on a user-role tool result, so + /// `Agent::run` reports [`RuntimeError::EmptyAssistantResponse`] for the + /// missing assistant message. That is bookkeeping about the wrong + /// question here, and this helper answers the right one instead: with the + /// expected new detail present the run succeeded, and without it the run + /// is reported as the missing terminal call it was. + pub async fn run_to_output( + &mut self, + content: impl Into>, + options: RunOptions, + spec: TerminalOutputSpec, + ) -> Result, RuntimeError> { + let tool_name = unique_tool_name(&spec.tool_name); + let keeps_tools = spec.keeps_tools; + let terminal_tool = TerminalOutputTool { + name: tool_name.clone(), + description: spec.description, + schema: spec.schema, + agent_id: self.id.clone(), + }; + self.runtime.register_scoped_tool(&self.id, terminal_tool); + *self + .terminal_tool_gate + .lock() + .expect("terminal tool gate poisoned") = Some(TerminalToolGate { + tool_name: tool_name.clone(), + keeps_tools, + }); + let _guard = TerminalToolGuard { + runtime: self.runtime.clone(), + agent_id: self.id.clone(), + tool_name: tool_name.clone(), + gate: Arc::clone(&self.terminal_tool_gate), + }; + + let run_result = self.run(content, options).await; + let terminal_result = self.terminal_result(&tool_name); + + match (run_result, terminal_result) { + (Ok(_), Some((details, message))) + | (Err(RuntimeError::EmptyAssistantResponse), Some((details, message))) => { + let value = serde_json::from_value(details).map_err(|error| { + RuntimeError::MalformedProviderEvent(format!( + "terminal output did not match the requested type: {error}" + )) + })?; + Ok(FinalOutput { value, message }) + } + (Ok(_) | Err(RuntimeError::EmptyAssistantResponse), None) => { + Err(RuntimeError::MalformedProviderEvent( + "run completed without invoking the expected terminal tool".to_string(), + )) + } + (Err(error), _) => Err(error), + } + } + + fn terminal_result(&self, tool_name: &str) -> Option<(Value, Message)> { + // Generated names include a per-call timestamp and counter, so scanning + // the whole transcript remains stale-safe even if auto-compaction + // replaced earlier items and changed every numeric index during the run. + let items = self.transcript().items(); + let expected_ids = items + .iter() + .filter_map(|item| item.message.as_ref()) + .filter(|message| message.role == Role::Assistant) + .flat_map(|message| message.content.iter()) + .filter_map(|block| match block { + ContentBlock::ToolUse { id, name, .. } if name == tool_name => Some(id.clone()), + _ => None, + }) + .collect::>(); + let last = items.last()?; + let message = last.message.clone()?; + let result_ids = message.content.iter().filter_map(|block| match block { + ContentBlock::ToolResult { tool_use_id, .. } => Some(tool_use_id), + _ => None, + }); + + for tool_use_id in result_ids { + if expected_ids.contains(tool_use_id) + && let Some(details) = last.detail(tool_use_id) + { + return Some((details.clone(), message)); + } + } + None + } +} + +struct TerminalOutputTool { + name: String, + description: String, + schema: Value, + agent_id: String, +} + +impl ToolDefinition for TerminalOutputTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder(self.name.clone()) + .description(self.description.clone()) + .input_schema(self.schema.clone()) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .terminal() + .build() + } +} + +#[async_trait] +impl ToolExecutor for TerminalOutputTool { + async fn execute_mut_output( + &self, + ctx: ToolContext<'_>, + input: Value, + ) -> Result { + if ctx.agent_id != self.agent_id { + return Err("terminal tool belongs to a different agent".to_string()); + } + Ok(ToolOutput::structured(input.clone()) + .with_details(input) + .terminating()) + } +} + +struct TerminalToolGuard { + runtime: RuntimeHandle, + agent_id: String, + tool_name: String, + gate: Arc>>, +} + +impl Drop for TerminalToolGuard { + fn drop(&mut self) { + let mut gate = self.gate.lock().expect("terminal tool gate poisoned"); + if gate + .as_ref() + .is_some_and(|open| open.tool_name == self.tool_name) + { + *gate = None; + } + drop(gate); + self.runtime + .unregister_scoped_tool(&self.agent_id, &self.tool_name); + } +} + +fn unique_tool_name(base: &str) -> String { + let mut base = base + .chars() + .map(|character| { + if character.is_ascii_alphanumeric() || character == '_' { + character + } else { + '_' + } + }) + .take(14) + .collect::(); + if base.is_empty() { + base = "output".to_string(); + } + let id = NEXT_TERMINAL_TOOL_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() as u64; + format!("mentra_terminal_{base}_{timestamp:016x}_{id:016x}") +} + +#[cfg(test)] +mod tests { + use super::unique_tool_name; + + #[test] + fn generated_tool_names_fit_common_provider_limits() { + let name = unique_tool_name("a name with punctuation and far too many characters"); + assert!(name.len() <= 64); + assert!( + name.chars() + .all(|character| character.is_ascii_alphanumeric() || character == '_') + ); + } +} diff --git a/vendor/mentra/src/agent/tests.rs b/vendor/mentra/src/agent/tests.rs new file mode 100644 index 0000000..c81a4af --- /dev/null +++ b/vendor/mentra/src/agent/tests.rs @@ -0,0 +1,16 @@ +mod budgets; +mod pending; +mod round_strategy; +mod runtime; +mod runtime_compact; +mod runtime_memory; +mod runtime_resume; +mod runtime_snapshot; +mod runtime_tasks; +mod runtime_tools; +mod runtime_volatile_store; +mod steering; +mod support; +mod terminal_output; +mod tool_output; +mod tool_paging; diff --git a/vendor/mentra/src/agent/tests/budgets.rs b/vendor/mentra/src/agent/tests/budgets.rs new file mode 100644 index 0000000..4177b6e --- /dev/null +++ b/vendor/mentra/src/agent/tests/budgets.rs @@ -0,0 +1,776 @@ +use crate::{ + BuiltinProvider, ContentBlock, Role, Runtime, TokenUsage, + agent::AgentEvent, + error::RuntimeError, + provider::{ContentBlockDelta, ContentBlockStart, ProviderEvent}, + runtime::{CancellationToken, EarlyEnd, RunOptions}, +}; + +use super::support::{ + ScriptedProvider, StaticTool, StopTrippingTool, StreamScript, model_info, ok_stream, +}; + +/// Builds a [`TokenUsage`] reporting only `input_tokens`/`output_tokens`, the two +/// fields [`RunOptions::token_budget`] is evaluated against. +fn usage(input_tokens: u64, output_tokens: u64) -> TokenUsage { + TokenUsage { + input_tokens: Some(input_tokens), + output_tokens: Some(output_tokens), + ..Default::default() + } +} + +/// Like `support::tool_use_stream`, but also reports `usage` via `MessageDelta` +/// so a round-boundary [`RunOptions::token_budget`] check has something to +/// evaluate. +fn tool_use_stream_with_usage( + model: &str, + id: &str, + name: &str, + input_json: &str, + usage: TokenUsage, +) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{id}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: id.to_string(), + name: name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageDelta { + stop_reason: None, + usage: Some(usage), + }, + ProviderEvent::MessageStopped, + ]) +} + +/// Like `support::text_stream`, but also reports `usage` via `MessageDelta`. +fn text_stream_with_usage(model: &str, text: &str, usage: TokenUsage) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageDelta { + stop_reason: None, + usage: Some(usage), + }, + ProviderEvent::MessageStopped, + ]) +} + +#[tokio::test] +async fn token_budget_stops_gracefully_after_the_round_that_crosses_it() { + // Round 1 reports usage that reaches the budget exactly; round 2 (a text + // response) must never be requested. Because the run stops before a final + // assistant message, `Agent::run` reports `EmptyAssistantResponse` — the same + // honest "stopped before a final answer" outcome `RunOptions::stop` produces + // at the identical boundary — while the gathered tool round stays committed + // rather than rolled back. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream_with_usage( + &model.id, + "call-1", + "probe_tool", + r#"{"value":"hi"}"#, + usage(60, 40), + ), + // Must never be requested: the budget trips before round 2 starts. + text_stream_with_usage(&model.id, "must not run", usage(1, 1)), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe_tool", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let result = agent + .run( + vec![ContentBlock::text("go")], + RunOptions { + token_budget: Some(100), + ..Default::default() + }, + ) + .await; + + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + assert_eq!( + agent.history().len(), + 3, + "the round that crossed the budget stays committed, not rolled back" + ); + assert_eq!( + provider_handle.recorded_requests().await.len(), + 1, + "the budget halted the run before a second model request" + ); +} + +#[tokio::test] +async fn absent_token_budget_ignores_reported_usage() { + // With `token_budget: None` (the default), no amount of reported usage stops + // the run early — the seam is inert, reproducing today's behavior exactly. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream_with_usage( + &model.id, + "call-1", + "probe_tool", + r#"{"value":"hi"}"#, + usage(10_000, 10_000), + ), + text_stream_with_usage(&model.id, "done", usage(10_000, 10_000)), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe_tool", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let message = agent + .run(vec![ContentBlock::text("go")], RunOptions::default()) + .await + .expect("run completes normally despite large reported usage"); + + assert_eq!(message.text(), "done"); + assert_eq!(provider_handle.recorded_requests().await.len(), 2); + assert_eq!(agent.history().len(), 4); +} + +#[tokio::test] +async fn a_crossed_budget_reports_why_the_turn_ended() { + // The defect this signal exists for: the run above ends correctly — the + // round that crossed the bound stays committed, nothing is rolled back — + // and says nothing about *why* it ended, since a run the model finished + // returns exactly the same way. `ended_early` is the runner's own answer, + // recorded at the boundary it refused to start another round at, and read + // back through a clone of the options the run was given. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream_with_usage( + &model.id, + "call-1", + "probe_tool", + r#"{"value":"hi"}"#, + usage(60, 40), + ), + text_stream_with_usage(&model.id, "must not run", usage(1, 1)), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe_tool", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let options = RunOptions { + token_budget: Some(100), + ..Default::default() + }; + let result = agent + .run(vec![ContentBlock::text("go")], options.clone()) + .await; + + assert_eq!( + options.ended_early(), + Some(EarlyEnd::TokenBudget), + "the run must report the bound that ended it, not leave it to be inferred" + ); + // The behavior around the report is unchanged, and pinned here beside it so + // that adding the signal cannot quietly turn a graceful end into a rollback. + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + assert_eq!( + agent.history().len(), + 3, + "the round that crossed the budget stays committed, not rolled back" + ); + assert_eq!(provider_handle.recorded_requests().await.len(), 1); +} + +#[tokio::test] +async fn a_requested_stop_is_not_reported_as_a_crossed_budget() { + // Both graceful bounds end a turn at the same boundary in the same way, so + // the signal is only worth anything if it tells them apart. The budget here + // is set and nowhere near crossed: what ends the turn is the tool tripping + // the stop token, and that is what must be reported. + let model = model_info("model", BuiltinProvider::Anthropic); + let stop = CancellationToken::default(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream_with_usage( + &model.id, + "call-1", + "stop_probe", + r#"{"value":"enough"}"#, + usage(10, 10), + ), + text_stream_with_usage(&model.id, "must not run", usage(1, 1)), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StopTrippingTool::new("stop_probe", stop.clone())) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let options = RunOptions { + stop: Some(stop), + token_budget: Some(10_000), + ..Default::default() + }; + let result = agent + .run(vec![ContentBlock::text("go")], options.clone()) + .await; + + assert_eq!(options.ended_early(), Some(EarlyEnd::StopRequested)); + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + assert_eq!( + options.reported_tokens(), + 20, + "the bound was never near: a caller recomputing it would have found nothing" + ); + assert_eq!(provider_handle.recorded_requests().await.len(), 1); +} + +#[tokio::test] +async fn a_stop_and_a_crossed_budget_together_report_the_stop() { + // The one round both spends the whole bound and trips the stop, so both + // conditions hold at the boundary that ends the turn. The stop wins: it is + // an instruction the caller issued, and the turn would have ended there + // with no budget set at all. Reporting the budget would tell a caller its + // allowance ran out when what happened is that it asked to stop. + let model = model_info("model", BuiltinProvider::Anthropic); + let stop = CancellationToken::default(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream_with_usage( + &model.id, + "call-1", + "stop_probe", + r#"{"value":"enough"}"#, + usage(60, 60), + ), + text_stream_with_usage(&model.id, "must not run", usage(1, 1)), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StopTrippingTool::new("stop_probe", stop.clone())) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let options = RunOptions { + stop: Some(stop), + token_budget: Some(100), + ..Default::default() + }; + let _ = agent + .run(vec![ContentBlock::text("go")], options.clone()) + .await; + + assert!( + options.reported_tokens() >= 100, + "the budget really is crossed, so the precedence is what decides the report" + ); + assert_eq!(options.ended_early(), Some(EarlyEnd::StopRequested)); +} + +#[tokio::test] +async fn a_turn_that_runs_to_completion_reports_nothing() { + // The default that keeps the signal honest: a bound that was set but never + // reached leaves nothing behind, so `Some(..)` always means the runner + // ended the turn rather than the model finishing it. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream_with_usage(&model.id, "done", usage(10, 10))], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let options = RunOptions { + stop: Some(CancellationToken::default()), + token_budget: Some(10_000), + ..Default::default() + }; + let message = agent + .run(vec![ContentBlock::text("go")], options.clone()) + .await + .expect("the run completes under both bounds"); + + assert_eq!(message.text(), "done"); + assert_eq!(options.ended_early(), None); +} + +#[tokio::test] +async fn a_turn_that_ends_on_the_budget_and_still_answers_reports_why() { + // The case that makes the signal load-bearing rather than convenient. The + // model finished its message *and* a steer was queued behind it, so the + // runner checks whether another request is available, finds the bound + // crossed, and returns the message it has. The turn is an ordinary `Ok` + // carrying an ordinary final answer — indistinguishable from a turn that + // ran to completion except through what the runner recorded, and the still + // pending steer is the work that got left behind. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream_with_usage(&model.id, "answered first", usage(60, 40)), + text_stream_with_usage(&model.id, "must not run", usage(1, 1)), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let steering = agent.steering_handle(); + steering.steer(vec![ContentBlock::text("and then this")]); + + let options = RunOptions { + token_budget: Some(100), + ..Default::default() + }; + let message = agent + .run(vec![ContentBlock::text("go")], options.clone()) + .await + .expect("a turn that ends on the budget after a committed message succeeds"); + + assert_eq!(message.text(), "answered first"); + assert_eq!( + options.ended_early(), + Some(EarlyEnd::TokenBudget), + "a successful turn is exactly where an unreported bound is invisible" + ); + assert!( + steering.has_pending(), + "the steer no request could carry is kept, not consumed" + ); + assert_eq!(provider_handle.recorded_requests().await.len(), 1); +} + +#[tokio::test] +async fn child_run_shares_cancellation_with_parent() { + // `RunOptions::child` carries the parent's `cancellation` token forward, so a + // parent cancel stops a child run threaded with the derived options — even + // though the two runs are on different agents and never call into each other. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream_with_usage( + &model.id, + "should not complete", + usage(1, 1), + )], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut child_agent = runtime.spawn("child", model).expect("spawn child agent"); + + let cancellation = CancellationToken::default(); + let parent_options = RunOptions { + cancellation: Some(cancellation.clone()), + ..Default::default() + }; + let child_options = parent_options.child(); + cancellation.cancel(); + + let error = child_agent + .run(vec![ContentBlock::text("go")], child_options) + .await + .expect_err("a cancelled parent token must stop the derived child run"); + + assert!(matches!(error, RuntimeError::Cancelled)); +} + +#[tokio::test] +async fn child_usage_counts_toward_shared_token_budget() { + // Parent and child share one token-accounting handle via `RunOptions::child`: + // neither run's own usage alone crosses the budget, but their combined total + // does, so the child's run stops gracefully at the shared bound. + let model = model_info("model", BuiltinProvider::Anthropic); + let parent_provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream_with_usage( + &model.id, + "parent done", + usage(40, 20), + )], + ); + let parent_runtime = Runtime::empty_builder() + .with_provider_instance(parent_provider) + .build() + .expect("build runtime"); + let mut parent_agent = parent_runtime + .spawn("parent", model.clone()) + .expect("spawn parent"); + + let parent_options = RunOptions { + token_budget: Some(100), + ..Default::default() + }; + parent_agent + .run(vec![ContentBlock::text("go")], parent_options.clone()) + .await + .expect("parent run completes under budget"); + assert_eq!( + parent_options.reported_tokens(), + 60, + "parent alone stays under the shared bound" + ); + + let child_options = parent_options.child(); + assert_eq!( + child_options.reported_tokens(), + 60, + "the derived child starts from the parent's already-reported usage" + ); + + let child_provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + // parent(60) + child(50) = 110, crossing the shared bound of 100. + tool_use_stream_with_usage( + &model.id, + "call-1", + "probe_tool", + r#"{"value":"hi"}"#, + usage(30, 20), + ), + text_stream_with_usage(&model.id, "must not run", usage(1, 1)), + ], + ); + let child_provider_handle = child_provider.clone(); + let child_runtime = Runtime::empty_builder() + .with_provider_instance(child_provider) + .with_tool(StaticTool::success("probe_tool", "ok")) + .build() + .expect("build runtime"); + let mut child_agent = child_runtime.spawn("child", model).expect("spawn child"); + + let result = child_agent + .run(vec![ContentBlock::text("go")], child_options) + .await; + + assert!( + matches!(result, Err(RuntimeError::EmptyAssistantResponse)), + "the child stops gracefully once the combined parent+child usage crosses the bound" + ); + assert_eq!( + child_provider_handle.recorded_requests().await.len(), + 1, + "the shared bound halted the child before its second round" + ); +} + +/// Drains the events an agent emitted during a finished run. +fn collect_events(receiver: &mut tokio::sync::broadcast::Receiver) -> Vec { + let mut events = Vec::new(); + while let Ok(event) = receiver.try_recv() { + events.push(event); + } + events +} + +/// The `input + output` totals of every [`AgentEvent::UsageReport`] in `events`, +/// in the order they were emitted. +fn reported_usage_totals(events: &[AgentEvent]) -> Vec { + events + .iter() + .filter_map(|event| match event { + AgentEvent::UsageReport { + input_tokens, + output_tokens, + .. + } => Some(input_tokens + output_tokens), + _ => None, + }) + .collect() +} + +#[tokio::test] +async fn delegated_subagent_usage_counts_against_the_parent_token_budget() { + // The `task` intrinsic is the one child run mentra drives itself: the model + // asks for it from inside the parent's run. Running it on the parent's + // `RunOptions::child` is what stops a model from delegating its way past the + // budget its own run was given — the child reports into the parent's + // accounting handle, and the parent ends at the next round boundary once the + // combined total crosses the bound. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + // Parent round 1 delegates, reporting 60 of the 100-token bound. + tool_use_stream_with_usage( + &model.id, + "parent-task", + "task", + r#"{"prompt":"delegate"}"#, + usage(40, 20), + ), + // The delegated run spends 50 more, taking the shared total to 110. + text_stream_with_usage(&model.id, "child summary", usage(30, 20)), + // Parent round 2 must never be requested. + text_stream_with_usage(&model.id, "parent done", usage(1, 1)), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let options = RunOptions { + token_budget: Some(100), + ..Default::default() + }; + let result = agent + .run(vec![ContentBlock::text("delegate that")], options.clone()) + .await; + + assert_eq!( + options.reported_tokens(), + 110, + "the delegated run must report into the parent's accounting handle, not a fresh one" + ); + assert!( + matches!(result, Err(RuntimeError::EmptyAssistantResponse)), + "the parent stops gracefully at the boundary after delegated spend crossed the bound, \ + reporting the same 'stopped before a final answer' outcome any tripped bound does" + ); + assert_eq!( + provider_handle.recorded_requests().await.len(), + 2, + "one parent round and one delegated round: the parent never got a second round" + ); +} + +#[tokio::test] +async fn parent_cancellation_reaches_the_delegated_subagent() { + // The delegated run shares the parent's cancellation token, so it ends at + // its own next round boundary rather than running on unreachable while the + // parent is torn down. `cancel_probe` trips the very token the parent's + // options hold, so the child honoring it can only mean it inherited that + // token rather than a default `None`. + let model = model_info("model", BuiltinProvider::Anthropic); + let cancellation = CancellationToken::default(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream_with_usage( + &model.id, + "parent-task", + "task", + r#"{"prompt":"delegate"}"#, + usage(1, 1), + ), + tool_use_stream_with_usage( + &model.id, + "child-tool", + "cancel_probe", + r#"{"value":"trip it"}"#, + usage(1, 1), + ), + // Never requested: the child checks the shared token before its + // second round. + text_stream_with_usage(&model.id, "child must not continue", usage(1, 1)), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_tool(StopTrippingTool::new("cancel_probe", cancellation.clone())) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let error = agent + .run( + vec![ContentBlock::text("delegate that")], + RunOptions { + cancellation: Some(cancellation), + ..Default::default() + }, + ) + .await + .expect_err("a cancelled run must fail rather than finish"); + + assert!(matches!(error, RuntimeError::Cancelled)); + assert_eq!( + provider_handle.recorded_requests().await.len(), + 2, + "the child stopped at its own round boundary; without the shared token it would \ + have run a second round before the parent ever saw the cancellation" + ); +} + +#[tokio::test] +async fn delegated_usage_reports_reach_the_parent_event_stream() { + // The accounting fix alone would leave a parent's observer blind to + // delegated spend, since a subagent has its own event bus. Relaying the + // child's `UsageReport` keeps a stream that sums usage agreeing with the + // shared handle the budget is checked against. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream_with_usage( + &model.id, + "parent-task", + "task", + r#"{"prompt":"delegate"}"#, + usage(40, 20), + ), + text_stream_with_usage(&model.id, "child summary", usage(30, 20)), + text_stream_with_usage(&model.id, "parent done", usage(5, 5)), + ], + ); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let mut events = agent.subscribe_events(); + + let options = RunOptions::default(); + let message = agent + .run(vec![ContentBlock::text("delegate that")], options.clone()) + .await + .expect("the run completes with no bound set"); + + assert_eq!(message.text(), "parent done"); + let totals = reported_usage_totals(&collect_events(&mut events)); + assert_eq!( + totals, + vec![60, 50, 10], + "the parent's stream carries the delegated round's usage between its own two rounds" + ); + assert_eq!( + totals.iter().sum::(), + options.reported_tokens(), + "what an observer sums from the stream matches what the budget is checked against" + ); +} + +#[tokio::test] +async fn delegating_with_the_budget_already_spent_fails_the_delegation() { + // A round is always allowed to finish, so the round that crosses the bound + // can still be the one asking to delegate. The child then inherits an + // already-exceeded budget and does zero rounds. That surfaces as a failed + // delegation the parent can see, not as a silent empty success — and the + // provider is never called on the child's behalf. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + // This single round spends the whole 100-token bound and delegates. + tool_use_stream_with_usage( + &model.id, + "parent-task", + "task", + r#"{"prompt":"delegate"}"#, + usage(60, 60), + ), + // Neither the child nor a second parent round may be requested. + text_stream_with_usage(&model.id, "must not run", usage(1, 1)), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let result = agent + .run( + vec![ContentBlock::text("delegate that")], + RunOptions { + token_budget: Some(100), + ..Default::default() + }, + ) + .await; + + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + assert_eq!( + provider_handle.recorded_requests().await.len(), + 1, + "the delegated run stopped at its first boundary without a model request" + ); + let subagents = agent.watch_snapshot().borrow().subagents.clone(); + assert_eq!(subagents.len(), 1); + assert!( + matches!( + &subagents[0].status, + crate::agent::SpawnedAgentStatus::Failed(message) + if message == "run completed without a final assistant message" + ), + "the exhausted delegation is recorded as failed, not finished: {:?}", + subagents[0].status + ); +} diff --git a/vendor/mentra/src/agent/tests/pending.rs b/vendor/mentra/src/agent/tests/pending.rs new file mode 100644 index 0000000..e148448 --- /dev/null +++ b/vendor/mentra/src/agent/tests/pending.rs @@ -0,0 +1,304 @@ +use serde_json::json; + +use crate::{ + ContentBlock, Message, Role, + agent::{AgentEvent, PendingAssistantTurn}, + provider::{ContentBlockDelta, ContentBlockStart, ProviderEvent, TokenUsage}, + tool::ToolCall, +}; + +#[test] +fn text_turn_commits_after_message_stop() { + let mut pending = PendingAssistantTurn::default(); + + assert!( + pending + .apply(ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: "model".to_string(), + role: Role::Assistant, + }) + .unwrap() + .is_empty() + ); + assert!( + pending + .apply(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }) + .unwrap() + .is_empty() + ); + + let derived = pending + .apply(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("Hello".to_string()), + }) + .unwrap(); + assert_eq!( + derived, + vec![AgentEvent::TextDelta { + delta: "Hello".to_string(), + full_text: "Hello".to_string(), + }] + ); + + pending + .apply(ProviderEvent::ContentBlockStopped { index: 0 }) + .unwrap(); + pending.apply(ProviderEvent::MessageStopped).unwrap(); + + assert_eq!( + pending.to_message().unwrap(), + Message::assistant(ContentBlock::text("Hello")) + ); +} + +#[test] +fn thinking_turn_emits_text_only_deltas_and_commits_signature_at_block_close() { + let provenance = crate::ReasoningProvenance { + provider: crate::ProviderId::new("anthropic-edge"), + model: "claude-test".to_string(), + format: crate::ReasoningFormat::AnthropicSigned, + }; + let mut pending = PendingAssistantTurn::default(); + pending + .apply(ProviderEvent::MessageStarted { + id: "msg-thinking".to_string(), + model: "claude-test".to_string(), + role: Role::Assistant, + }) + .unwrap(); + pending + .apply(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Thinking { + encrypted_content: None, + id: None, + provenance: Some(provenance.clone()), + redacted: false, + }, + }) + .unwrap(); + + assert_eq!( + pending + .apply(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingText("private ".to_string()), + }) + .unwrap(), + vec![AgentEvent::ReasoningDelta { + delta: "private ".to_string(), + full_text: "private ".to_string(), + }] + ); + assert_eq!( + pending + .apply(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingText("chain".to_string()), + }) + .unwrap(), + vec![AgentEvent::ReasoningDelta { + delta: "chain".to_string(), + full_text: "private chain".to_string(), + }] + ); + assert!( + pending + .apply(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingSignature("opaque-signature".to_string()), + }) + .unwrap() + .is_empty() + ); + pending + .apply(ProviderEvent::ContentBlockStopped { index: 0 }) + .unwrap(); + pending.apply(ProviderEvent::MessageStopped).unwrap(); + + assert_eq!( + pending.to_message().unwrap(), + Message::assistant(ContentBlock::Thinking { + thinking: "private chain".to_string(), + signature: Some("opaque-signature".to_string()), + encrypted_content: None, + id: None, + provenance: Some(provenance), + redacted: false, + }) + ); +} + +#[test] +fn tool_use_turn_emits_ready_event_and_parses_call() { + let mut pending = PendingAssistantTurn::default(); + + pending + .apply(ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: "model".to_string(), + role: Role::Assistant, + }) + .unwrap(); + pending + .apply(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "tool-1".to_string(), + name: "echo_tool".to_string(), + }, + }) + .unwrap(); + pending + .apply(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(r#"{"value":"hi"}"#.to_string()), + }) + .unwrap(); + + let derived = pending + .apply(ProviderEvent::ContentBlockStopped { index: 0 }) + .unwrap(); + assert_eq!( + derived, + vec![AgentEvent::ToolUseReady { + index: 0, + call: ToolCall { + id: "tool-1".to_string(), + name: "echo_tool".to_string(), + input: json!({ "value": "hi" }), + }, + }] + ); + + pending.apply(ProviderEvent::MessageStopped).unwrap(); + assert_eq!(pending.ready_tool_calls().unwrap().len(), 1); +} + +#[test] +fn pending_turn_rejects_missing_stop_and_recovers_from_malformed_tool_json() { + let mut text_pending = PendingAssistantTurn::default(); + text_pending + .apply(ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: "model".to_string(), + role: Role::Assistant, + }) + .unwrap(); + text_pending + .apply(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }) + .unwrap(); + text_pending + .apply(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("Hello".to_string()), + }) + .unwrap(); + text_pending + .apply(ProviderEvent::ContentBlockStopped { index: 0 }) + .unwrap(); + assert!(text_pending.to_message().is_err()); + + let mut tool_pending = PendingAssistantTurn::default(); + tool_pending + .apply(ProviderEvent::MessageStarted { + id: "msg-2".to_string(), + model: "model".to_string(), + role: Role::Assistant, + }) + .unwrap(); + tool_pending + .apply(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "tool-1".to_string(), + name: "broken_tool".to_string(), + }, + }) + .unwrap(); + tool_pending + .apply(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson("{".to_string()), + }) + .unwrap(); + assert!( + tool_pending + .apply(ProviderEvent::ContentBlockStopped { index: 0 }) + .unwrap() + .is_empty() + ); + tool_pending.apply(ProviderEvent::MessageStopped).unwrap(); + + assert!(tool_pending.ready_tool_calls().unwrap().is_empty()); + assert_eq!(tool_pending.invalid_tool_uses().len(), 1); + assert_eq!( + tool_pending.to_message().unwrap(), + Message { + role: Role::Assistant, + content: Vec::new(), + } + ); +} + +#[test] +fn pending_turn_tracks_latest_usage_without_affecting_message() { + let mut pending = PendingAssistantTurn::default(); + pending + .apply(ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: "model".to_string(), + role: Role::Assistant, + }) + .unwrap(); + pending + .apply(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }) + .unwrap(); + pending + .apply(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("Hello".to_string()), + }) + .unwrap(); + pending + .apply(ProviderEvent::ContentBlockStopped { index: 0 }) + .unwrap(); + pending + .apply(ProviderEvent::MessageDelta { + stop_reason: Some("stop".to_string()), + usage: Some(TokenUsage { + input_tokens: Some(12), + output_tokens: Some(3), + total_tokens: Some(15), + ..TokenUsage::default() + }), + }) + .unwrap(); + pending.apply(ProviderEvent::MessageStopped).unwrap(); + + assert_eq!( + pending.usage(), + Some(&TokenUsage { + input_tokens: Some(12), + output_tokens: Some(3), + total_tokens: Some(15), + ..TokenUsage::default() + }) + ); + assert_eq!(pending.stop_reason(), Some("stop")); + assert_eq!( + pending.to_message().unwrap(), + Message::assistant(ContentBlock::text("Hello")) + ); +} diff --git a/vendor/mentra/src/agent/tests/round_strategy.rs b/vendor/mentra/src/agent/tests/round_strategy.rs new file mode 100644 index 0000000..2f0a1fe --- /dev/null +++ b/vendor/mentra/src/agent/tests/round_strategy.rs @@ -0,0 +1,500 @@ +use std::{ + collections::VecDeque, + sync::{Arc, Mutex}, +}; + +use async_trait::async_trait; + +use crate::{ + BuiltinProvider, ContentBlock, Message, ReasoningEffort, ReasoningOptions, Role, Runtime, + agent::{ + ReasoningChange, RoundAdjustment, RoundBoundary, RoundContext, RoundDecision, + RoundStrategy, RoundToolResult, + }, + error::RuntimeError, + runtime::RunOptions, +}; + +use super::support::{ScriptedProvider, StaticTool, model_info, text_stream, tool_use_stream}; + +/// One boundary observation recorded by [`ScriptedStrategy`]. +#[derive(Clone)] +struct Observation { + boundary: RoundBoundary, + rounds_completed: usize, + model_requests: usize, + assistant_text: Option, + tool_results: Vec, +} + +/// A scripted decision returned at a boundary. An exhausted script yields +/// [`RoundDecision::proceed`]. +enum DecisionScript { + Inject(String), + Stop, + Switch(RoundAdjustment), +} + +/// A [`RoundStrategy`] that records every boundary it observes and replays a +/// scripted sequence of decisions. +struct ScriptedStrategy { + log: Mutex>, + decisions: Mutex>, +} + +impl ScriptedStrategy { + fn new(decisions: Vec) -> Arc { + Arc::new(Self { + log: Mutex::new(Vec::new()), + decisions: Mutex::new(decisions.into()), + }) + } + + fn observations(&self) -> Vec { + self.log.lock().expect("strategy log poisoned").clone() + } + + fn invocation_count(&self) -> usize { + self.log.lock().expect("strategy log poisoned").len() + } +} + +#[async_trait] +impl RoundStrategy for ScriptedStrategy { + async fn on_round(&self, ctx: RoundContext<'_>) -> RoundDecision { + self.log + .lock() + .expect("strategy log poisoned") + .push(Observation { + boundary: ctx.boundary(), + rounds_completed: ctx.rounds_completed(), + model_requests: ctx.model_requests(), + assistant_text: ctx.assistant_message().map(Message::text), + tool_results: ctx.tool_results().to_vec(), + }); + match self + .decisions + .lock() + .expect("strategy decisions poisoned") + .pop_front() + { + None => RoundDecision::proceed(), + Some(DecisionScript::Inject(text)) => { + RoundDecision::inject(vec![ContentBlock::text(text)]) + } + Some(DecisionScript::Stop) => RoundDecision::stop(), + Some(DecisionScript::Switch(adjust)) => RoundDecision::Continue(adjust), + } + } +} + +/// Captured outcome of a two-round probe session (tool round then text round). +struct SessionCapture { + history: Vec, + request_models: Vec, + request_messages: Vec>, +} + +/// Runs the canonical two-round probe session (one tool round, one terminal text +/// round), optionally attaching a proceed-everywhere strategy built from +/// `decisions`, and captures the transcript and recorded requests. +async fn run_probe_session(decisions: Option>) -> SessionCapture { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "probe_tool", r#"{"value":"hi"}"#), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe_tool", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let options = match decisions { + Some(decisions) => { + RunOptions::default().with_round_strategy(ScriptedStrategy::new(decisions)) + } + None => RunOptions::default(), + }; + agent + .run(vec![ContentBlock::text("hi")], options) + .await + .expect("run succeeds"); + + let requests = provider_handle.recorded_requests().await; + SessionCapture { + history: agent.history().to_vec(), + request_models: requests.iter().map(|r| r.model.to_string()).collect(), + request_messages: requests.iter().map(|r| r.messages.to_vec()).collect(), + } +} + +#[tokio::test] +async fn strategy_observes_both_round_boundaries_in_order() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "probe_tool", r#"{"value":"hi"}"#), + text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe_tool", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let strategy = ScriptedStrategy::new(vec![]); + agent + .run( + vec![ContentBlock::text("hi")], + RunOptions::default().with_round_strategy(strategy.clone()), + ) + .await + .expect("run succeeds"); + + let observations = strategy.observations(); + assert_eq!(observations.len(), 2); + + // Boundary (a): fired after the committed tool round, before the next round. + assert_eq!( + observations[0].boundary, + RoundBoundary::ToolResultsCommitted + ); + assert_eq!(observations[0].rounds_completed, 1); + assert_eq!(observations[0].model_requests, 1); + assert!(observations[0].assistant_text.is_none()); + assert_eq!(observations[0].tool_results.len(), 1); + assert_eq!(observations[0].tool_results[0].tool_use_id, "call-1"); + assert_eq!(observations[0].tool_results[0].tool_name, "probe_tool"); + assert!(!observations[0].tool_results[0].is_error); + + // Boundary (b): fired after the committed tool-free assistant message. + assert_eq!( + observations[1].boundary, + RoundBoundary::AssistantMessageCommitted + ); + assert_eq!(observations[1].rounds_completed, 2); + assert_eq!(observations[1].model_requests, 2); + assert_eq!(observations[1].assistant_text.as_deref(), Some("done")); + assert!(observations[1].tool_results.is_empty()); +} + +#[tokio::test] +async fn none_strategy_matches_continue_strategy_byte_identical() { + // The `None` default and a proceed-everywhere strategy must produce an + // identical transcript and identical recorded requests: the seam is inert. + let baseline = run_probe_session(None).await; + let with_strategy = run_probe_session(Some(vec![])).await; + + assert_eq!(baseline.request_models.len(), 2, "two rounds ran"); + assert_eq!(baseline.history, with_strategy.history); + assert_eq!(baseline.request_models, with_strategy.request_models); + assert_eq!(baseline.request_messages, with_strategy.request_messages); +} + +#[tokio::test] +async fn injected_context_reaches_next_provider_request() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first"), + text_stream(&model.id, "second"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let strategy = ScriptedStrategy::new(vec![DecisionScript::Inject( + "please call finish_investigation".to_string(), + )]); + let message = agent + .run( + vec![ContentBlock::text("hi")], + RunOptions::default().with_round_strategy(strategy), + ) + .await + .expect("run succeeds"); + + assert_eq!(message.text(), "second"); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 2, "the injection forced a second round"); + let second_round_has_injection = requests[1] + .messages + .iter() + .any(|m| m.role == Role::User && m.text().contains("please call finish_investigation")); + assert!( + second_round_has_injection, + "injected corrective context must reach the next provider request" + ); +} + +#[tokio::test] +async fn model_and_reasoning_switch_applies_to_next_round() { + let model_a = model_info("model-a", BuiltinProvider::Anthropic); + let model_b = model_info("model-b", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model_a.clone(), model_b.clone()], + vec![ + tool_use_stream(&model_a.id, "call-1", "probe_tool", r#"{"value":"hi"}"#), + text_stream(&model_b.id, "done"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe_tool", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model_a).expect("spawn agent"); + + let adjust = RoundAdjustment::new() + .with_model(model_b) + .with_reasoning(ReasoningChange::Set(ReasoningOptions { + effort: Some(ReasoningEffort::High), + summary: None, + })); + let strategy = ScriptedStrategy::new(vec![DecisionScript::Switch(adjust)]); + agent + .run( + vec![ContentBlock::text("hi")], + RunOptions::default().with_round_strategy(strategy), + ) + .await + .expect("run succeeds"); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 2); + // Round 1 (before the switch) uses the original model and no reasoning. + assert_eq!(requests[0].model.as_ref(), "model-a"); + assert_eq!(requests[0].provider_request_options.reasoning, None); + // Round 2 (after the switch) uses the switched model and reasoning. + assert_eq!(requests[1].model.as_ref(), "model-b"); + assert_eq!( + requests[1].provider_request_options.reasoning, + Some(ReasoningOptions { + effort: Some(ReasoningEffort::High), + summary: None, + }) + ); +} + +#[tokio::test] +async fn stop_after_tool_round_commits_transcript_and_halts() { + // A Stop returned at the tool-round boundary matches `RunOptions::stop`: the + // gathered transcript is committed (not rolled back) and no further model + // request is made. Because the last committed message is a tool result rather + // than an assistant message, `Agent::run` surfaces `EmptyAssistantResponse` — + // the honest "stopped before a final answer" outcome, identical to + // `RunOptions::stop` firing at the same boundary. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "probe_tool", r#"{"value":"hi"}"#), + text_stream(&model.id, "must not run"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe_tool", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let strategy = ScriptedStrategy::new(vec![DecisionScript::Stop]); + let result = agent + .run( + vec![ContentBlock::text("go")], + RunOptions::default().with_round_strategy(strategy), + ) + .await; + + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + assert_eq!( + agent.history().len(), + 3, + "the gathered tool round is committed, not rolled back" + ); + assert_eq!( + provider_handle.recorded_requests().await.len(), + 1, + "the stop halted the run before a second model request" + ); +} + +#[tokio::test] +async fn stop_at_assistant_boundary_returns_committed_message() { + // At the assistant boundary a Stop commits the transcript and returns Ok with + // the terminal message (the finish_run path, not rollback). + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "final answer")], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let strategy = ScriptedStrategy::new(vec![DecisionScript::Stop]); + let message = agent + .run( + vec![ContentBlock::text("hi")], + RunOptions::default().with_round_strategy(strategy), + ) + .await + .expect("a stop at the assistant boundary returns Ok"); + + assert_eq!(message.text(), "final answer"); + assert_eq!( + agent.history().len(), + 2, + "user + committed assistant message" + ); +} + +#[tokio::test] +async fn strategy_state_does_not_outlive_run() { + // Two sequential runs on one runtime/agent, each with its own strategy + // instance, share nothing: each strategy observes only its own run's boundary. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first"), + text_stream(&model.id, "second"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let strategy_one = ScriptedStrategy::new(vec![]); + agent + .run( + vec![ContentBlock::text("run one")], + RunOptions::default().with_round_strategy(strategy_one.clone()), + ) + .await + .expect("first run succeeds"); + + let strategy_two = ScriptedStrategy::new(vec![]); + agent + .run( + vec![ContentBlock::text("run two")], + RunOptions::default().with_round_strategy(strategy_two.clone()), + ) + .await + .expect("second run succeeds"); + + assert_eq!( + strategy_one.invocation_count(), + 1, + "the first strategy saw only its own run" + ); + assert_eq!( + strategy_two.invocation_count(), + 1, + "the second strategy saw only its own run" + ); +} + +#[tokio::test] +async fn assistant_boundary_continue_returns_inject_runs_another_round() { + // Continue at the assistant boundary accepts the terminal message and returns. + { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "solo")], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let strategy = ScriptedStrategy::new(vec![]); + let message = agent + .run( + vec![ContentBlock::text("hi")], + RunOptions::default().with_round_strategy(strategy), + ) + .await + .expect("run succeeds"); + + assert_eq!(message.text(), "solo"); + assert_eq!( + provider_handle.recorded_requests().await.len(), + 1, + "continue returns without another round" + ); + } + + // Inject at the assistant boundary prevents returning and runs another round. + { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first"), + text_stream(&model.id, "second"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let strategy = + ScriptedStrategy::new(vec![DecisionScript::Inject("keep going".to_string())]); + let message = agent + .run( + vec![ContentBlock::text("hi")], + RunOptions::default().with_round_strategy(strategy), + ) + .await + .expect("run succeeds"); + + assert_eq!( + message.text(), + "second", + "the injected round produced the terminal message" + ); + assert_eq!( + provider_handle.recorded_requests().await.len(), + 2, + "inject forced another model round" + ); + } +} diff --git a/vendor/mentra/src/agent/tests/runtime.rs b/vendor/mentra/src/agent/tests/runtime.rs new file mode 100644 index 0000000..0236c08 --- /dev/null +++ b/vendor/mentra/src/agent/tests/runtime.rs @@ -0,0 +1,999 @@ +use std::{ + sync::{Arc, Mutex}, + time::Duration, +}; + +use async_trait::async_trait; +use serde_json::{Value, json}; +use tokio::time::sleep; + +use crate::{ + BuiltinProvider, ContentBlock, Role, + provider::{ContentBlockDelta, ContentBlockStart, ProviderError, ProviderEvent, TokenUsage}, + runtime::{ + RunOptions, Runtime, RuntimeHook, RuntimeHookEvent, RuntimePolicy, ShellValidationMode, + is_transient_runtime_error, + }, + tool::{ + ToolAuthorizationDecision, ToolAuthorizationOutcome, ToolAuthorizationRequest, + ToolAuthorizer, ToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSpec, + }, +}; + +use super::support::{ScriptedProvider, StaticTool, erroring_stream, model_info, ok_stream}; + +#[tokio::test] +async fn run_respects_tool_budget() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![tool_use_stream( + &model.id, + "tool-1", + "test_tool", + r#"{"value":"hi"}"#, + )], + ); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("test_tool", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + let error = agent + .run( + vec![ContentBlock::Text { + text: "hi".to_string(), + }], + RunOptions { + tool_budget: Some(0), + ..RunOptions::default() + }, + ) + .await + .expect_err("tool budget should abort run"); + + assert!(matches!( + error, + crate::runtime::RuntimeError::ToolBudgetExceeded(0) + )); +} + +#[tokio::test] +async fn workspace_bounded_policy_keeps_local_shell_disabled() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "pwd", "shell", r#"{"command":"pwd"}"#), + text_stream("done"), + ], + ); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_policy(RuntimePolicy::workspace_bounded( + std::env::current_dir().expect("current directory"), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "run pwd".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: true, .. } + if content.contains("disabled by the runtime policy") + )); +} + +#[tokio::test] +async fn enforced_shell_validation_blocks_destructive_command_and_emits_hook() { + let model = model_info("model", BuiltinProvider::Anthropic); + let sentinel = std::env::temp_dir().join(format!( + "mentra-shell-validation-{}-{}", + std::process::id(), + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("system time") + .as_nanos() + )); + std::fs::create_dir_all(&sentinel).expect("create sentinel directory"); + let input = json!({ "command": format!("rm -rf {}", sentinel.display()) }).to_string(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "rm", "shell", &input), + text_stream("done"), + ], + ); + let recorded = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_policy( + RuntimePolicy::read_only(std::env::current_dir().expect("current directory")) + // This test exercises validation before execution. The + // destructive command is expected to stop at that boundary. + .allow_shell_commands(true) + .shell_validation(ShellValidationMode::Enforce), + ) + .with_hook(RecordingHook { + events: recorded.clone(), + }) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "delete everything".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: true, .. } + if content.contains("not allowed in read-only mode") + )); + assert!( + recorded + .lock() + .expect("hook events poisoned") + .iter() + .any(|event| matches!( + event, + RuntimeHookEvent::AuthorizationDenied { action, detail, .. } + if action == "shell_validation" && detail.contains("not allowed") + )) + ); + assert!( + sentinel.exists(), + "blocked command must not reach the executor" + ); + std::fs::remove_dir_all(sentinel).expect("remove sentinel directory"); +} + +#[tokio::test] +async fn shell_authorization_preview_carries_validation_intent() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "rm", + "shell", + r#"{"command":"rm -rf /tmp/mentra-preview-never-execute"}"#, + ), + text_stream("done"), + ], + ); + let requests = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_policy(RuntimePolicy::workspace_bounded( + std::env::current_dir().expect("current directory"), + )) + .with_tool_authorizer(RecordingAuthorizer::deny( + "destructive command denied", + requests.clone(), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "classify this command".to_string(), + }]) + .await + .expect("send"); + + let requests = requests.lock().expect("requests poisoned"); + let validation = &requests[0].preview.structured_input["validation"]; + assert_eq!(validation["mode"], "off"); + assert_eq!(validation["intent"], "destructive"); + assert_eq!(validation["outcome"], "prompt"); + assert!(validation["reason"].as_str().is_some()); +} + +#[tokio::test] +async fn tool_authorizer_allows_tool_execution_and_captures_preview() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-1", "test_tool", r#"{"value":"hi"}"#), + text_stream("done"), + ], + ); + let requests = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("test_tool", "ok")) + .with_tool_authorizer(RecordingAuthorizer::allow(requests.clone())) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hi".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: false, .. } + if content.to_display_string() == "ok" + )); + + let requests = requests.lock().expect("requests poisoned"); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].tool_name, "test_tool"); + assert_eq!( + requests[0].preview.structured_input, + json!({ "value": "hi" }) + ); +} + +#[tokio::test] +async fn tool_authorizer_can_prompt_shell_tool_and_emit_hooks() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "tool-shell", + "shell", + r#"{"command":"python -c 'print(1)'","justification":"needed for validation"}"#, + ), + text_stream("done"), + ], + ); + let requests = Arc::new(Mutex::new(Vec::new())); + let recorded = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_policy(RuntimePolicy::default().allow_shell_commands(true)) + .with_tool_authorizer(RecordingAuthorizer::prompt( + "needs manual review", + requests.clone(), + )) + .with_hook(RecordingHook { + events: recorded.clone(), + }) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "run python".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: true, .. } + if content.contains("Tool execution requires approval: needs manual review") + )); + + let requests = requests.lock().expect("requests poisoned"); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].tool_name, "shell"); + assert_eq!( + requests[0].preview.structured_input["kind"].as_str(), + Some("shell") + ); + + let events = recorded.lock().expect("hook events poisoned").clone(); + assert!(events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ToolAuthorizationStarted { tool_name, .. } if tool_name == "shell" + ))); + assert!(events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ToolAuthorizationFinished { tool_name, outcome, .. } + if tool_name == "shell" && *outcome == ToolAuthorizationOutcome::Prompt + ))); + assert!(events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ToolAuthorizationBlocked { tool_name, outcome, .. } + if tool_name == "shell" && *outcome == ToolAuthorizationOutcome::Prompt + ))); + assert!(!events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ToolExecutionStarted { tool_name, .. } if tool_name == "shell" + ))); +} + +#[tokio::test] +async fn tool_authorizer_can_deny_background_run() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "tool-bg", + "background_run", + r#"{"command":"python -c 'print(1)'","justification":"background probe"}"#, + ), + text_stream("done"), + ], + ); + let requests = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_policy( + RuntimePolicy::default() + .allow_shell_commands(true) + .allow_background_commands(true), + ) + .with_tool_authorizer(RecordingAuthorizer::deny( + "background tasks require review", + requests.clone(), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "run background work".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: true, .. } + if content.contains("Tool execution denied: background tasks require review") + )); + + let requests = requests.lock().expect("requests poisoned"); + assert_eq!( + requests[0].preview.structured_input["kind"].as_str(), + Some("background_run") + ); + assert_eq!( + requests[0].preview.structured_input["background"].as_bool(), + Some(true) + ); +} + +#[tokio::test] +async fn tool_authorizer_errors_block_execution() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-1", "test_tool", r#"{"value":"hi"}"#), + text_stream("done"), + ], + ); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("test_tool", "ok")) + .with_tool_authorizer(RecordingAuthorizer::error("authorizer unavailable")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hi".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: true, .. } + if content.contains("authorizer unavailable") + )); +} + +#[tokio::test] +async fn tool_authorizer_timeout_blocks_execution() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-1", "test_tool", r#"{"value":"hi"}"#), + text_stream("done"), + ], + ); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("test_tool", "ok")) + .with_tool_authorizer(RecordingAuthorizer::delayed_allow( + Duration::from_millis(50), + Duration::from_millis(10), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hi".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: true, .. } + if content.contains("authorizer timed out") + )); +} + +#[tokio::test] +async fn files_tool_authorization_preview_exposes_resolved_paths() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-read", + "files", + r#"{"operations":[{"op":"read","path":"README.md","offset":1,"limit":1}]}"#, + ), + text_stream("done"), + ], + ); + let requests = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_tool_authorizer(RecordingAuthorizer::allow(requests.clone())) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "read the readme".to_string(), + }]) + .await + .expect("send"); + + let requests = requests.lock().expect("requests poisoned"); + assert_eq!(requests.len(), 1); + let operations = requests[0].preview.structured_input["operations"] + .as_array() + .expect("operations array"); + assert_eq!(operations[0]["op"].as_str(), Some("read")); + assert!( + operations[0]["resolved_path"] + .as_str() + .expect("resolved path") + .ends_with("README.md") + ); + assert!(operations[0].get("content").is_none()); +} + +#[tokio::test] +async fn custom_hooks_observe_model_and_tool_execution() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-1", "test_tool", r#"{"value":"hi"}"#), + text_stream("done"), + ], + ); + let recorded = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("test_tool", "ok")) + .with_hook(RecordingHook { + events: recorded.clone(), + }) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hi".to_string(), + }]) + .await + .expect("send"); + + let events = recorded.lock().expect("hook events poisoned").clone(); + assert!( + events + .iter() + .any(|event| matches!(event, RuntimeHookEvent::ModelRequestStarted { .. })) + ); + assert!(events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ModelResponseFinished { + success: true, + stop_reason: Some(reason), + usage: None, + .. + } if reason == "tool_use" + ))); + assert!(events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ToolExecutionStarted { tool_name, .. } if tool_name == "test_tool" + ))); + assert!(events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ToolExecutionFinished { tool_name, is_error: false, .. } + if tool_name == "test_tool" + ))); +} + +#[tokio::test] +async fn tools_can_read_registered_app_context() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-ctx", "app_context_tool", r#"{}"#), + text_stream("done"), + ], + ); + let app_state = Arc::new(TestAppState { + label: "configured", + }); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_context(app_state.clone()) + .with_tool(AppContextTool) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + assert_eq!( + runtime + .app_context::() + .expect("app context should be registered") + .label, + "configured" + ); + + agent + .send(vec![ContentBlock::Text { + text: "use the app_context_tool".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: false, .. } + if content.to_display_string() == "configured" + )); +} + +#[tokio::test] +async fn tool_execution_timeout_returns_tool_error() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-slow", "slow_tool", r#"{}"#), + text_stream("done"), + ], + ); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(SlowTool) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "run the slow tool".to_string(), + }]) + .await + .expect("send"); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: true, .. } + if content.contains("timed out after 20ms") + )); +} + +#[tokio::test] +async fn model_response_finished_hook_reports_usage_after_successful_commit() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-usage".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("done".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageDelta { + stop_reason: Some("end_turn".to_string()), + usage: Some(TokenUsage { + input_tokens: Some(12), + output_tokens: Some(5), + total_tokens: Some(17), + ..TokenUsage::default() + }), + }, + ProviderEvent::MessageStopped, + ])], + ); + let recorded = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_hook(RecordingHook { + events: recorded.clone(), + }) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hi".to_string(), + }]) + .await + .expect("send"); + + let events = recorded.lock().expect("hook events poisoned").clone(); + assert!(events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ModelResponseFinished { + success: true, + stop_reason: Some(reason), + usage: Some(TokenUsage { + input_tokens: Some(12), + output_tokens: Some(5), + total_tokens: Some(17), + .. + }), + .. + } if reason == "end_turn" + ))); +} + +#[tokio::test] +async fn model_response_finished_hook_reports_stream_failures_without_usage() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![erroring_stream( + vec![ + ProviderEvent::MessageStarted { + id: "msg-fail".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("par".to_string()), + }, + ], + ProviderError::MalformedStream("boom".to_string()), + )], + ); + let recorded = Arc::new(Mutex::new(Vec::new())); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_hook(RecordingHook { + events: recorded.clone(), + }) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + let error = agent + .send(vec![ContentBlock::Text { + text: "hi".to_string(), + }]) + .await + .expect_err("send should fail"); + assert!(matches!( + error, + crate::runtime::RuntimeError::FailedToStreamResponse(ProviderError::MalformedStream(_)) + )); + + let events = recorded.lock().expect("hook events poisoned").clone(); + assert!(events.iter().any(|event| matches!( + event, + RuntimeHookEvent::ModelResponseFinished { + success: false, + usage: None, + error: Some(message), + .. + } if message.contains("malformed provider stream: boom") + ))); +} + +#[test] +fn transient_runtime_error_helper_matches_provider_retry_policy() { + let transient = crate::runtime::RuntimeError::FailedToStreamResponse(ProviderError::Http { + status: reqwest::StatusCode::TOO_MANY_REQUESTS, + body: String::new(), + retry_after: None, + }); + let permanent = crate::runtime::RuntimeError::FailedToStreamResponse( + ProviderError::InvalidRequest("bad request".to_string()), + ); + + assert!(is_transient_runtime_error(&transient)); + assert!(!is_transient_runtime_error(&permanent)); + assert!(!is_transient_runtime_error( + &crate::runtime::RuntimeError::EmptyAssistantResponse + )); +} + +#[derive(Debug)] +struct TestAppState { + label: &'static str, +} + +struct AppContextTool; + +#[async_trait] +impl ToolDefinition for AppContextTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("app_context_tool") + .description("Return a value from the runtime app context.") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .build() + } +} + +#[async_trait] +impl ToolExecutor for AppContextTool { + async fn execute_mut(&self, ctx: ToolContext<'_>, _input: Value) -> ToolResult { + Ok(ctx.app_context::()?.label.to_string()) + } +} + +struct SlowTool; + +#[async_trait] +impl ToolDefinition for SlowTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("slow_tool") + .description("Sleep long enough to trigger a timeout.") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .execution_timeout(Duration::from_millis(20)) + .build() + } +} + +#[async_trait] +impl ToolExecutor for SlowTool { + async fn execute_mut(&self, _ctx: ToolContext<'_>, _input: Value) -> ToolResult { + sleep(Duration::from_millis(60)).await; + Ok("finished".to_string()) + } +} + +#[derive(Clone)] +struct RecordingHook { + events: Arc>>, +} + +impl RuntimeHook for RecordingHook { + fn on_event( + &self, + _store: &dyn crate::runtime::AuditStore, + event: &RuntimeHookEvent, + ) -> Result<(), crate::runtime::RuntimeError> { + self.events + .lock() + .expect("hook events poisoned") + .push(event.clone()); + Ok(()) + } +} + +enum AuthorizerBehavior { + Allow, + Prompt(String), + Deny(String), + Error(String), + DelayedAllow(Duration), +} + +struct RecordingAuthorizer { + behavior: AuthorizerBehavior, + requests: Arc>>, + timeout: Option, +} + +impl RecordingAuthorizer { + fn allow(requests: Arc>>) -> Self { + Self { + behavior: AuthorizerBehavior::Allow, + requests, + timeout: None, + } + } + + fn prompt(reason: &str, requests: Arc>>) -> Self { + Self { + behavior: AuthorizerBehavior::Prompt(reason.to_string()), + requests, + timeout: None, + } + } + + fn error(reason: &str) -> Self { + Self { + behavior: AuthorizerBehavior::Error(reason.to_string()), + requests: Arc::new(Mutex::new(Vec::new())), + timeout: None, + } + } + + fn deny(reason: &str, requests: Arc>>) -> Self { + Self { + behavior: AuthorizerBehavior::Deny(reason.to_string()), + requests, + timeout: None, + } + } + + fn delayed_allow(delay: Duration, timeout: Duration) -> Self { + Self { + behavior: AuthorizerBehavior::DelayedAllow(delay), + requests: Arc::new(Mutex::new(Vec::new())), + timeout: Some(timeout), + } + } +} + +#[async_trait] +impl ToolAuthorizer for RecordingAuthorizer { + async fn authorize( + &self, + request: &ToolAuthorizationRequest, + ) -> Result { + self.requests + .lock() + .expect("requests poisoned") + .push(request.clone()); + + match &self.behavior { + AuthorizerBehavior::Allow => Ok(ToolAuthorizationDecision::allow()), + AuthorizerBehavior::Prompt(reason) => { + Ok(ToolAuthorizationDecision::prompt(reason.clone())) + } + AuthorizerBehavior::Deny(reason) => Ok(ToolAuthorizationDecision::deny(reason.clone())), + AuthorizerBehavior::Error(reason) => { + Err(crate::runtime::RuntimeError::Store(reason.clone())) + } + AuthorizerBehavior::DelayedAllow(delay) => { + sleep(*delay).await; + Ok(ToolAuthorizationDecision::allow()) + } + } + } + + fn timeout(&self) -> Option { + self.timeout + } +} + +fn text_stream(text: &str) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-text".to_string(), + model: "model".to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageDelta { + stop_reason: Some("end_turn".to_string()), + usage: None, + }, + ProviderEvent::MessageStopped, + ]) +} + +fn tool_use_stream( + model: &str, + id: &str, + name: &str, + input_json: &str, +) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{id}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: id.to_string(), + name: name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageDelta { + stop_reason: Some("tool_use".to_string()), + usage: None, + }, + ProviderEvent::MessageStopped, + ]) +} diff --git a/vendor/mentra/src/agent/tests/runtime_compact.rs b/vendor/mentra/src/agent/tests/runtime_compact.rs new file mode 100644 index 0000000..aa91ae0 --- /dev/null +++ b/vendor/mentra/src/agent/tests/runtime_compact.rs @@ -0,0 +1,1127 @@ +use std::{ + fs, + path::PathBuf, + sync::atomic::{AtomicU64, Ordering}, + time::{SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::{ + BuiltinProvider, ContentBlock, Message, Role, TranscriptKind, + agent::{AgentConfig, AgentEvent, CompactionConfig, CompactionTrigger}, + compaction::{CompactionExecutionMode, CompactionMode}, + provider::{CompactionInputItem, CompactionResponse, ProviderCapabilities, Request}, + runtime::{Runtime, SqliteRuntimeStore}, + tool::{ + ToolContext, ToolDefinition, ToolDurability, ToolExecutor, ToolOutput, ToolSideEffectLevel, + ToolSpec, + }, +}; + +use crate::provider::ProviderError; + +use super::support::{ + ScriptedProvider, SessionGenerator, StaticTool, erroring_stream, model_info, text_stream, + tool_use_stream, +}; + +/// A tool whose output is long enough to trigger micro-compaction's +/// content-collapse threshold, and which attaches opaque `details` — used to +/// prove micro-compaction (a request-projection concern) never touches the +/// canonical transcript item's metadata (M3 test 5), and that full +/// compaction preserves it on items outside the compacted prefix. +struct DetailsTool { + output: String, +} + +impl DetailsTool { + fn new(output: impl Into) -> Self { + Self { + output: output.into(), + } + } +} + +#[async_trait] +impl ToolDefinition for DetailsTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("details_tool") + .description("test tool: returns long output plus opaque details") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .build() + } +} + +#[async_trait] +impl ToolExecutor for DetailsTool { + async fn execute_mut_output( + &self, + _ctx: ToolContext<'_>, + _input: Value, + ) -> Result { + Ok(ToolOutput::text(self.output.clone()).with_details(json!({ "marker": "keep-me" }))) + } +} + +#[tokio::test] +async fn micro_compaction_only_rewrites_old_tool_results_in_requests() { + let model = model_info("model", BuiltinProvider::Anthropic); + let long_output = "x".repeat(140); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-1", "echo_tool", r#"{"value":"one"}"#), + tool_use_stream(&model.id, "tool-2", "echo_tool", r#"{"value":"two"}"#), + tool_use_stream(&model.id, "tool-3", "echo_tool", r#"{"value":"three"}"#), + tool_use_stream(&model.id, "tool-4", "echo_tool", r#"{"value":"four"}"#), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("echo_tool", &long_output)) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + keep_recent_tool_results: 2, + auto_compact_threshold_tokens: None, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await + .unwrap(); + + assert_eq!( + agent.history()[2], + Message::user(ContentBlock::ToolResult { + tool_use_id: "tool-1".to_string(), + content: long_output.clone().into(), + is_error: false, + }) + ); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 5); + let final_tool_results = tool_result_contents(&requests[4]); + assert_eq!( + final_tool_results, + vec![ + "[Previous: used echo_tool]".to_string(), + "[Previous: used echo_tool]".to_string(), + long_output.clone(), + long_output, + ] + ); +} + +// M3 test 5: micro-compaction rewrites old tool results only in the +// *request projection* (`micro_compacted_history`, a fresh clone of the +// transcript's `Message`s built on every `stream_turn`) — it never touches +// the canonical `TranscriptItem`s themselves, so the details a host attached +// to the collapsed call survive on the stored item even though the outgoing +// request no longer carries the original content. +#[tokio::test] +async fn micro_compaction_leaves_stored_item_details_intact() { + let model = model_info("model", BuiltinProvider::Anthropic); + let long_output = "x".repeat(140); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-1", "details_tool", r#"{}"#), + tool_use_stream(&model.id, "tool-2", "details_tool", r#"{}"#), + tool_use_stream(&model.id, "tool-3", "details_tool", r#"{}"#), + tool_use_stream(&model.id, "tool-4", "details_tool", r#"{}"#), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(DetailsTool::new(long_output.clone())) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + keep_recent_tool_results: 2, + auto_compact_threshold_tokens: None, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await + .unwrap(); + + // The outgoing request for the final round collapsed the two oldest + // tool results (tool-1, tool-2) to a placeholder — proof micro-compaction + // actually ran. + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 5); + let final_tool_results = tool_result_contents(&requests[4]); + assert_eq!( + final_tool_results, + vec![ + "[Previous: used details_tool]".to_string(), + "[Previous: used details_tool]".to_string(), + long_output.clone(), + long_output, + ] + ); + + // The canonical transcript item for the collapsed tool-1 call still + // carries its full details, untouched by the request-side collapse. + let item = agent + .transcript() + .items() + .iter() + .find(|item| { + matches!( + item.kind, + TranscriptKind::ToolExchange { + tool_use_id: Some(ref id), + .. + } if id == "tool-1" + ) + }) + .expect("tool-1's transcript item is still present"); + assert_eq!(item.detail("tool-1"), Some(&json!({ "marker": "keep-me" }))); +} + +// M3 (spec A6 / plan M5 prerequisite, not the M5 guarantee itself): a +// compacted-away prefix is summarized away, but the tail — the last +// assistant tool_use + its tool result — is copied into the replacement +// transcript verbatim (`replacement.extend_from_slice`, +// `src/compaction.rs::StandardCompactionEngine::compact`), so a +// details-bearing item in that tail keeps its metadata "for free" through +// `TranscriptItem`'s derived `Clone`. This is a cheap sanity check on that +// existing copy path, not the exhaustive metadata-preservation contract +// (ADR-0001 §6), which is a separate, later slice. +#[tokio::test] +async fn auto_compaction_preserves_details_on_tail_items() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "details_tool", r#"{}"#), + text_stream(&model.id, "summary"), + text_stream(&model.id, "after compact"), + ], + ); + let transcript_dir = temp_dir("details-preserve-compact"); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(DetailsTool::new("tool result")) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(1), + transcript_dir, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "run the details tool".to_string(), + }]) + .await + .unwrap(); + + let compaction = collect_events(&mut events) + .into_iter() + .find_map(|event| match event { + AgentEvent::ContextCompacted { details } => Some(details), + _ => None, + }) + .expect("expected a compaction event"); + assert_eq!(compaction.trigger, CompactionTrigger::Auto); + assert_eq!( + compaction.preserved_items, 2, + "the assistant tool_use and its tool result stay in the tail" + ); + + let item = agent + .transcript() + .items() + .iter() + .find(|item| matches!(item.kind, TranscriptKind::ToolExchange { .. })) + .expect("compaction preserved the tool exchange item in the tail"); + assert_eq!(item.detail("call-1"), Some(&json!({ "marker": "keep-me" }))); +} + +#[tokio::test] +async fn auto_compaction_persists_transcript_and_rewrites_history() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first done"), + text_stream(&model.id, "summary"), + text_stream(&model.id, "second done"), + ], + ); + let provider_handle = provider.clone(); + let transcript_dir = temp_dir("auto-compact"); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(1), + transcript_dir: transcript_dir.clone(), + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "first".to_string(), + }]) + .await + .unwrap(); + agent + .send(vec![ContentBlock::Text { + text: "second".to_string(), + }]) + .await + .unwrap(); + + assert_eq!(agent.history().len(), 4); + assert_eq!(agent.history()[0].role, Role::User); + assert_eq!(message_text(&agent.history()[0]), "first"); + assert!(message_text(&agent.history()[1]).contains("[Compaction summary]")); + assert!(message_text(&agent.history()[1]).contains("Progress: summary")); + + let transcripts = fs::read_dir(&transcript_dir) + .expect("read transcript dir") + .map(|entry| entry.expect("read transcript entry").path()) + .collect::>(); + assert_eq!(transcripts.len(), 1); + + let transcript = fs::read_to_string(&transcripts[0]).expect("read transcript"); + assert_eq!(transcript.lines().count(), 3); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 3); + assert!(requests[1].tools.is_empty()); + assert_eq!(requests[1].tool_choice, None); + assert_eq!(message_text(&requests[2].messages[0]), "first"); + assert!( + requests[2] + .messages + .iter() + .any(|message| message_text(message).contains("Progress: summary")) + ); + + let compaction = collect_events(&mut events) + .into_iter() + .find_map(|event| match event { + AgentEvent::ContextCompacted { details } => Some(details), + _ => None, + }) + .expect("expected compaction event"); + assert_eq!(compaction.trigger, CompactionTrigger::Auto); + assert_eq!(compaction.replaced_items, 2); + assert_eq!(compaction.preserved_items, 1); + assert_eq!(compaction.preserved_user_turns, 1); + assert_eq!(compaction.preserved_delegation_results, 0); + assert_eq!(compaction.resulting_transcript_len, 3); + assert!(compaction.transcript_path.starts_with(&transcript_dir)); +} + +#[tokio::test] +async fn compact_tool_compacts_history_and_continues() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "compact-1", "compact", "{}"), + text_stream(&model.id, "summary"), + text_stream(&model.id, "after compact"), + ], + ); + let provider_handle = provider.clone(); + let transcript_dir = temp_dir("manual-compact"); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: None, + transcript_dir, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "please compact".to_string(), + }]) + .await + .unwrap(); + + assert_eq!(agent.history().len(), 5); + assert_eq!(message_text(&agent.history()[0]), "please compact"); + assert!(message_text(&agent.history()[1]).contains("[Compaction summary]")); + assert!(message_text(&agent.history()[1]).contains("Progress: summary")); + assert!(matches!( + &agent.history()[3].content[0], + ContentBlock::ToolResult { is_error: false, content, .. } + if content.starts_with("Context compacted. Transcript saved to ") + )); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 3); + assert!(requests[1].tools.is_empty()); + assert_eq!(requests[1].tool_choice, None); + assert_eq!(message_text(&requests[2].messages[0]), "please compact"); + assert!( + requests[2] + .messages + .iter() + .any(|message| message_text(message).contains("Progress: summary")) + ); + assert!(tool_names(&requests[0]).contains("compact")); + + let compaction = collect_events(&mut events) + .into_iter() + .find_map(|event| match event { + AgentEvent::ContextCompacted { details } => Some(details), + _ => None, + }) + .expect("expected compaction event"); + assert_eq!(compaction.trigger, CompactionTrigger::Manual); + assert_eq!(compaction.replaced_items, 1); + assert_eq!(compaction.preserved_items, 1); + assert_eq!(compaction.preserved_user_turns, 1); + assert_eq!(compaction.preserved_delegation_results, 0); + assert_eq!(compaction.resulting_transcript_len, 3); +} + +#[tokio::test] +async fn auto_compaction_degrades_gracefully_on_failure() { + let model = model_info("model", BuiltinProvider::Anthropic); + // Queue: first send response, then 3 retryable errors for compaction attempts, + // then the second send response. The compaction will fail all 3 attempts and + // degrade gracefully, allowing the second send to succeed. + let retryable_error = || { + erroring_stream( + vec![], + ProviderError::Retryable { + message: "rate limited".into(), + delay: None, + }, + ) + }; + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first done"), + retryable_error(), + retryable_error(), + retryable_error(), + text_stream(&model.id, "second done"), + ], + ); + let events_receiver = { + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(1), + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "first".to_string(), + }]) + .await + .unwrap(); + + // Second send triggers auto_compact_if_needed which fails all 3 attempts, + // then degrades gracefully, and the actual send succeeds. + agent + .send(vec![ContentBlock::Text { + text: "second".to_string(), + }]) + .await + .expect("second send must succeed despite compaction failures"); + + // History should have all 4 turns (no compaction was applied). + assert_eq!(agent.history().len(), 4, "history should have 4 items"); + + collect_events(&mut events) + }; + + // Should have seen 2 RetryAttempt events (attempts 1 and 2; attempt 3 exhausts + // without emitting because there is no further retry after the last attempt). + let retry_events: Vec<_> = events_receiver + .iter() + .filter(|e| matches!(e, AgentEvent::RetryAttempt { .. })) + .collect(); + assert_eq!( + retry_events.len(), + 2, + "expected 2 retry attempt events, got {}", + retry_events.len() + ); + + // No ContextCompacted event should have been emitted. + let compacted = events_receiver + .iter() + .any(|e| matches!(e, AgentEvent::ContextCompacted { .. })); + assert!(!compacted, "expected no ContextCompacted event"); +} + +#[tokio::test] +async fn remote_compaction_succeeds_when_provider_supports_it() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first done"), + text_stream(&model.id, "second done"), + ], + ) + .with_capabilities(ProviderCapabilities { + supports_history_compaction: true, + ..Default::default() + }); + + provider + .push_compact_response(Ok(CompactionResponse { + output: vec![CompactionInputItem::CompactionSummary { + content: "Summary of previous work".to_string(), + }], + })) + .await; + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(1), + mode: CompactionMode::PreferRemote, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "first".to_string(), + }]) + .await + .unwrap(); + agent + .send(vec![ContentBlock::Text { + text: "second".to_string(), + }]) + .await + .unwrap(); + + let compaction = collect_events(&mut events) + .into_iter() + .find_map(|event| match event { + AgentEvent::ContextCompacted { details } => Some(details), + _ => None, + }) + .expect("expected compaction event"); + assert_eq!(compaction.mode, CompactionExecutionMode::Remote); +} + +#[tokio::test] +async fn remote_compaction_falls_back_to_local_on_unsupported() { + let model = model_info("model", BuiltinProvider::Anthropic); + // Provider advertises remote support but compact() returns UnsupportedCapability + // (no compact scripts pushed — default error). + // Local summarization calls provider.stream(), so we need an extra text stream for it. + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first done"), + text_stream(&model.id, "summary"), + text_stream(&model.id, "second done"), + ], + ) + .with_capabilities(ProviderCapabilities { + supports_history_compaction: true, + ..Default::default() + }); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(1), + mode: CompactionMode::PreferRemote, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "first".to_string(), + }]) + .await + .unwrap(); + agent + .send(vec![ContentBlock::Text { + text: "second".to_string(), + }]) + .await + .unwrap(); + + let compaction = collect_events(&mut events) + .into_iter() + .find_map(|event| match event { + AgentEvent::ContextCompacted { details } => Some(details), + _ => None, + }) + .expect("expected compaction event"); + assert_eq!(compaction.mode, CompactionExecutionMode::Local); +} + +#[tokio::test] +async fn remote_compaction_falls_back_to_local_on_empty_remote_response() { + let model = model_info("model", BuiltinProvider::Anthropic); + // Provider advertises remote support but returns an empty response — compact_remotely + // returns Ok(None) which triggers a local fallback. + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first done"), + text_stream(&model.id, "summary"), + text_stream(&model.id, "second done"), + ], + ) + .with_capabilities(ProviderCapabilities { + supports_history_compaction: true, + ..Default::default() + }); + + provider + .push_compact_response(Ok(CompactionResponse { output: vec![] })) + .await; + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(1), + mode: CompactionMode::PreferRemote, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "first".to_string(), + }]) + .await + .unwrap(); + agent + .send(vec![ContentBlock::Text { + text: "second".to_string(), + }]) + .await + .unwrap(); + + let compaction = collect_events(&mut events) + .into_iter() + .find_map(|event| match event { + AgentEvent::ContextCompacted { details } => Some(details), + _ => None, + }) + .expect("expected compaction event"); + assert_eq!(compaction.mode, CompactionExecutionMode::Local); +} + +// --------------------------------------------------------------------------- +// Moderate CI integration tests — multi-turn sessions with compaction cycles +// --------------------------------------------------------------------------- + +/// Runs 50 turns with a low auto-compact threshold to trigger multiple compaction +/// cycles, then asserts that compaction fired at least twice and that history was +/// meaningfully reduced. +#[tokio::test] +async fn fifty_turn_session_with_multiple_compaction_cycles() { + let model = model_info("model", BuiltinProvider::Anthropic); + let transcript_dir = temp_dir("fifty-turn-multi-compact"); + + // Generate 300 scripted responses for 50 actual sends. + // With a very low threshold every turn can trigger compaction, and each + // compaction consumes one extra response for the local summarizer. + // 300 gives generous headroom even if compaction fires on every turn. + let scripts = SessionGenerator::new(&model.id) + .with_response_size(500) + .add_text_turns(300) + .build(); + + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], scripts); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + // threshold=1 guarantees compaction fires before every turn + // after the first response is committed to history. + auto_compact_threshold_tokens: Some(1), + transcript_dir, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + + let mut events = agent.subscribe_events(); + let mut compaction_count = 0usize; + + for i in 0..50u32 { + agent + .send(vec![ContentBlock::Text { + text: format!("Turn {i}"), + }]) + .await + .unwrap_or_else(|e| panic!("turn {i} failed: {e}")); + // Drain the event channel after each turn to avoid broadcast overflow. + compaction_count += collect_events(&mut events) + .iter() + .filter(|e| matches!(e, AgentEvent::ContextCompacted { .. })) + .count(); + } + + assert!( + compaction_count >= 2, + "expected at least 2 compaction cycles after 50 turns, got {compaction_count}" + ); + + // History should be compressed — without compaction it would be 100 messages + // (50 user + 50 assistant). With compaction each cycle replaces most history + // with a single summary message. + assert!( + agent.history().len() < 100, + "expected history to be compacted (< 100 messages), got {}", + agent.history().len() + ); +} + +/// Tests that a session survives persist → drop → rebuild → resume across a +/// compaction boundary. The resumed agent should be able to continue sending turns. +#[tokio::test] +async fn resumed_session_continues_after_compaction() { + let model = model_info("model", BuiltinProvider::Anthropic); + let transcript_dir = temp_dir("resume-after-compact"); + + // Use a persistent store so we can reopen it after dropping the runtime. + let store = temp_sqlite_store("resume-after-compact"); + + // Phase 1 — run 15 turns with a low threshold to ensure at least one + // compaction fires, then drop agent + runtime to persist state. + { + // With threshold=1, compaction fires on every turn after the first. + // 15 turns need 15 turn responses + 14 summarizer responses = 29 total. + // Generate 50 to give generous headroom. + let scripts = SessionGenerator::new(&model.id) + .with_response_size(500) + .add_text_turns(50) + .build(); + + let provider = + ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], scripts); + + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + + let mut agent = runtime + .spawn_with_config( + "agent", + model.clone(), + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(1), + transcript_dir, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + + let mut events = agent.subscribe_events(); + let mut compaction_count = 0usize; + + for i in 0..15u32 { + agent + .send(vec![ContentBlock::Text { + text: format!("Phase-1 turn {i}"), + }]) + .await + .unwrap_or_else(|e| panic!("phase-1 turn {i} failed: {e}")); + // Drain the event channel after each turn to avoid broadcast overflow. + compaction_count += collect_events(&mut events) + .iter() + .filter(|e| matches!(e, AgentEvent::ContextCompacted { .. })) + .count(); + } + assert!( + compaction_count >= 1, + "expected at least 1 compaction in phase 1, got {compaction_count}" + ); + + // Dropping agent then runtime persists state and releases the lease. + drop(agent); + drop(runtime); + } + + // Clear leases so the second runtime can acquire the agent. + clear_sqlite_leases(&store); + + // Phase 2 — rebuild runtime with the same store and resume the agent. + { + let scripts = SessionGenerator::new(&model.id) + .with_response_size(200) + .add_text_turns(10) + .build(); + + let provider = + ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], scripts); + + let new_runtime = Runtime::empty_builder() + .with_store(store) + .with_provider_instance(provider) + .build() + .expect("rebuild runtime"); + + let resumed_agents = new_runtime.resume_all().expect("resume_all"); + assert_eq!( + resumed_agents.len(), + 1, + "expected exactly one resumed agent" + ); + let mut agent = resumed_agents.into_iter().next().unwrap(); + + for i in 0..5u32 { + agent + .send(vec![ContentBlock::Text { + text: format!("Phase-2 turn {i}"), + }]) + .await + .unwrap_or_else(|e| panic!("phase-2 turn {i} failed: {e}")); + } + + // Resumed agent should have produced at least the 5 post-resume replies. + assert!( + !agent.history().is_empty(), + "resumed agent should have history after additional sends" + ); + } +} + +/// Smoke test: multiple compaction cycles must not panic or corrupt the session. +/// Verifies that history survives three or more compaction cycles over 30 turns. +#[tokio::test] +async fn compaction_chain_preserves_context_across_cycles() { + let model = model_info("model", BuiltinProvider::Anthropic); + let transcript_dir = temp_dir("compact-chain"); + + // With threshold=1, compaction fires on every turn after the first. + // 30 turns require 30 turn responses + 29 summarizer responses = 59 total. + // Generate 100 to give generous headroom. + let scripts = SessionGenerator::new(&model.id) + .with_response_size(500) + .add_text_turns(100) + .build(); + + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], scripts); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(1), + transcript_dir, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + + let mut events = agent.subscribe_events(); + let mut compaction_count = 0usize; + + for i in 0..30u32 { + agent + .send(vec![ContentBlock::Text { + text: format!("Turn {i}"), + }]) + .await + .unwrap_or_else(|e| panic!("turn {i} failed: {e}")); + // Drain the event channel after each turn to avoid broadcast overflow. + compaction_count += collect_events(&mut events) + .iter() + .filter(|e| matches!(e, AgentEvent::ContextCompacted { .. })) + .count(); + } + + assert!( + compaction_count >= 2, + "expected at least 2 compaction cycles after 30 turns, got {compaction_count}" + ); + + // After multiple compaction cycles the session must still be usable — + // history is non-empty and we didn't panic. + assert!( + !agent.history().is_empty(), + "history must not be empty after compaction chain" + ); +} + +fn temp_sqlite_store(label: &str) -> SqliteRuntimeStore { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = std::env::temp_dir().join(format!( + "mentra-runtime-compact-{label}-{timestamp}-{unique}.sqlite" + )); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).expect("create temp dir"); + } + SqliteRuntimeStore::new(path) +} + +fn clear_sqlite_leases(store: &SqliteRuntimeStore) { + let conn = rusqlite::Connection::open(store.path()).expect("open store"); + conn.execute("DELETE FROM leases", []) + .expect("clear leases"); +} + +fn tool_result_contents(request: &Request<'_>) -> Vec { + request + .messages + .iter() + .flat_map(|message| message.content.iter()) + .filter_map(|block| match block { + ContentBlock::ToolResult { content, .. } => Some(content.to_display_string()), + _ => None, + }) + .collect() +} + +fn tool_names(request: &Request<'_>) -> std::collections::HashSet { + request.tools.iter().map(|tool| tool.name.clone()).collect() +} + +fn message_text(message: &Message) -> &str { + message + .content + .iter() + .find_map(|block| match block { + ContentBlock::Text { text } => Some(text.as_str()), + _ => None, + }) + .unwrap_or("") +} + +fn collect_events(receiver: &mut tokio::sync::broadcast::Receiver) -> Vec { + let mut events = Vec::new(); + while let Ok(event) = receiver.try_recv() { + events.push(event); + } + events +} + +#[tokio::test] +async fn transcript_cleanup_prunes_old_files() { + use crate::compaction::cleanup_old_transcripts; + + let dir = temp_dir("cleanup-prune"); + + // Write 5 fake .jsonl files with increasing timestamps so sort order is deterministic. + let mut filenames = Vec::new(); + for i in 0..5u64 { + let name = format!("{:020}.jsonl", i); + let path = dir.join(&name); + fs::write(&path, b"{}").expect("write fake transcript"); + filenames.push(name); + } + + // Keep only 3 (the 2 oldest should be removed). + cleanup_old_transcripts(&dir, 3) + .await + .expect("cleanup should succeed"); + + let remaining: std::collections::BTreeSet = fs::read_dir(&dir) + .expect("read dir") + .map(|e| { + e.expect("dir entry") + .file_name() + .to_string_lossy() + .into_owned() + }) + .collect(); + + assert_eq!(remaining.len(), 3, "expected 3 files, got {remaining:?}"); + // The 3 newest files (indices 2, 3, 4) must survive. + for i in 2..5u64 { + let expected = format!("{:020}.jsonl", i); + assert!( + remaining.contains(&expected), + "expected {expected} to remain, got {remaining:?}" + ); + } + // The 2 oldest files (indices 0, 1) must be gone. + for i in 0..2u64 { + let expected = format!("{:020}.jsonl", i); + assert!( + !remaining.contains(&expected), + "expected {expected} to be deleted, got {remaining:?}" + ); + } +} + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +fn temp_dir(label: &str) -> PathBuf { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = std::env::temp_dir().join(format!( + "mentra-runtime-compact-{label}-{timestamp}-{unique}" + )); + fs::create_dir_all(&path).expect("create temp dir"); + path +} diff --git a/vendor/mentra/src/agent/tests/runtime_memory.rs b/vendor/mentra/src/agent/tests/runtime_memory.rs new file mode 100644 index 0000000..0326eb7 --- /dev/null +++ b/vendor/mentra/src/agent/tests/runtime_memory.rs @@ -0,0 +1,483 @@ +use std::{path::PathBuf, time::Duration}; + +use crate::{ + BuiltinProvider, ContentBlock, Message, Role, + agent::{AgentConfig, CompactionConfig, MemoryConfig}, + memory::{MemoryRecord, MemoryRecordKind, MemoryStore}, + provider::{ContentBlockDelta, ContentBlockStart, ProviderEvent}, + runtime::{HybridRuntimeStore, Runtime, SqliteRuntimeStore}, +}; + +use super::support::{ScriptedProvider, StreamScript, model_info, ok_stream}; + +#[tokio::test] +async fn automatic_memory_search_injects_recalled_context_without_persisting_it() { + let store = test_store("recalled-memory"); + store + .upsert_records(&[MemoryRecord { + record_id: "summary:agent:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Summary, + content: "The user prefers keeping memory automatic and bounded.".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("seed".to_string()), + pinned: false, + score: None, + }]) + .expect("seed records"); + + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "done")], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let agent_id = agent.id().to_string(); + store + .upsert_records(&[MemoryRecord { + record_id: format!("summary:{agent_id}:1"), + agent_id: agent_id.clone(), + kind: MemoryRecordKind::Summary, + content: "The user prefers keeping memory automatic and bounded.".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("seed".to_string()), + pinned: false, + score: None, + }]) + .expect("seed agent record"); + + agent + .send(vec![ContentBlock::Text { + text: "Help me design memory".to_string(), + }]) + .await + .expect("run"); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 1); + assert!(requests[0].messages.iter().any(|message| { + message_text(message).contains("") + && message_text(message).contains("memory automatic and bounded") + })); + assert!( + agent + .history() + .iter() + .all(|message| { !message_text(message).contains("") }) + ); +} + +#[tokio::test] +async fn successful_runs_are_ingested_and_searchable() { + let store = test_store("memory-ingest"); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "finished task")], + ); + + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let agent_id = agent.id().to_string(); + + agent + .send(vec![ContentBlock::Text { + text: "remember this plan".to_string(), + }]) + .await + .expect("run"); + + let records = wait_for_records(&store, &agent_id, "remember", 1).await; + assert_eq!(records.len(), 1); + assert_eq!(records[0].kind, MemoryRecordKind::Episode); + assert!(records[0].content.contains("remember this plan")); + assert!(records[0].content.contains("finished task")); +} + +#[tokio::test] +async fn sqlite_memory_search_is_namespaced_per_agent() { + let store = test_store("memory-isolation"); + store + .upsert_records(&[ + MemoryRecord { + record_id: "episode:a:1".to_string(), + agent_id: "agent-a".to_string(), + kind: MemoryRecordKind::Episode, + content: "shared phrase alpha".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("seed".to_string()), + pinned: false, + score: None, + }, + MemoryRecord { + record_id: "episode:b:1".to_string(), + agent_id: "agent-b".to_string(), + kind: MemoryRecordKind::Episode, + content: "shared phrase alpha".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("seed".to_string()), + pinned: false, + score: None, + }, + ]) + .expect("seed records"); + + let agent_a = store + .search_records("agent-a", "shared alpha", 10) + .expect("search agent a"); + let agent_b = store + .search_records("agent-b", "shared alpha", 10) + .expect("search agent b"); + + assert_eq!(agent_a.len(), 1); + assert_eq!(agent_b.len(), 1); + assert_eq!(agent_a[0].agent_id, "agent-a"); + assert_eq!(agent_b[0].agent_id, "agent-b"); +} + +#[tokio::test] +async fn compacted_summaries_are_searchable() { + let store = test_store("memory-compact"); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "compact-1", "compact", "{}"), + text_stream(&model.id, "summary about architecture"), + text_stream(&model.id, "after compact"), + ], + ); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + auto_compact_threshold_tokens: None, + transcript_dir: temp_dir("searchable-compact"), + ..Default::default() + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + let agent_id = agent.id().to_string(); + + agent + .send(vec![ContentBlock::Text { + text: "please compact".to_string(), + }]) + .await + .expect("run"); + + let records = store + .search_records(&agent_id, "architecture", 10) + .expect("search summaries"); + assert!(records.iter().any(|record| { + record.kind == MemoryRecordKind::Summary + && record.content.contains("summary about architecture") + })); +} + +#[tokio::test] +async fn hybrid_memory_recall_includes_provenance_in_hidden_context() { + let store = hybrid_store("recalled-memory-provenance"); + let model = model_info("model", BuiltinProvider::Anthropic); + store + .upsert_records(&[MemoryRecord { + record_id: "fact:agent:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Fact, + content: "The user prefers concise memory summaries.".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("manual_pin".to_string()), + pinned: true, + score: None, + }]) + .expect("seed records"); + + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "done")], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let agent_id = agent.id().to_string(); + store + .upsert_records(&[MemoryRecord { + record_id: format!("fact:{agent_id}:1"), + agent_id: agent_id.clone(), + kind: MemoryRecordKind::Fact, + content: "The user prefers concise memory summaries.".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("manual_pin".to_string()), + pinned: true, + score: None, + }]) + .expect("seed agent record"); + + agent + .send(vec![ContentBlock::Text { + text: "Design the memory flow".to_string(), + }]) + .await + .expect("run"); + + let requests = provider_handle.recorded_requests().await; + let injected = requests[0] + .messages + .iter() + .find_map(|message| { + let text = message_text(message); + text.contains("") + .then_some(text.to_string()) + }) + .expect("recalled memory"); + assert!(injected.contains("source=manual_pin")); + assert!(injected.contains("why=")); +} + +#[tokio::test] +async fn memory_pin_tool_creates_searchable_hybrid_memory() { + let store = hybrid_store("memory-pin-tool"); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "memory-pin-1", + "memory_pin", + r#"{"content":"Remember that the user likes short answers."}"#, + ), + text_stream(&model.id, "pinned"), + ], + ); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + memory: MemoryConfig { + write_tools_enabled: true, + ..Default::default() + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + let agent_id = agent.id().to_string(); + + agent + .send(vec![ContentBlock::Text { + text: "Please remember this.".to_string(), + }]) + .await + .expect("run"); + + let records = wait_for_records(&store, &agent_id, "short answers", 1).await; + assert!(records.iter().any(|record| { + record.kind == MemoryRecordKind::Fact + && record.pinned + && record.source.as_deref() == Some("manual_pin") + })); +} + +#[tokio::test] +async fn memory_forget_tool_hides_pinned_hybrid_memory() { + let store = hybrid_store("memory-forget-tool"); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "memory-forget-1", + "memory_forget", + r#"{"record_id":"fact:forget:1"}"#, + ), + text_stream(&model.id, "forgot"), + ], + ); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + memory: MemoryConfig { + write_tools_enabled: true, + ..Default::default() + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + let agent_id = agent.id().to_string(); + store + .upsert_records(&[MemoryRecord { + record_id: "fact:forget:1".to_string(), + agent_id: agent_id.clone(), + kind: MemoryRecordKind::Fact, + content: "The user likes short answers.".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("manual_pin".to_string()), + pinned: true, + score: None, + }]) + .expect("seed records"); + + agent + .send(vec![ContentBlock::Text { + text: "Forget that note.".to_string(), + }]) + .await + .expect("run"); + + let records = store + .search_records(&agent_id, "short answers", 10) + .expect("search records"); + assert!(records.is_empty()); +} + +async fn wait_for_records( + store: &impl MemoryStore, + agent_id: &str, + query: &str, + expected: usize, +) -> Vec { + for _ in 0..50 { + let records = store + .search_records(agent_id, query, 10) + .expect("search records"); + if records.len() >= expected { + return records; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + store + .search_records(agent_id, query, 10) + .expect("final search") +} + +fn test_store(prefix: &str) -> SqliteRuntimeStore { + SqliteRuntimeStore::new(temp_dir(prefix).join("runtime.sqlite")) +} + +fn hybrid_store(prefix: &str) -> HybridRuntimeStore { + let dir = temp_dir(prefix); + HybridRuntimeStore::with_memory_path(dir.join("runtime.sqlite"), dir.join("memory.sqlite")) +} + +fn temp_dir(prefix: &str) -> PathBuf { + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("time") + .as_nanos(); + std::env::temp_dir().join(format!("mentra-{prefix}-{nanos}")) +} + +fn text_stream(model: &str, text: &str) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn tool_use_stream(model: &str, id: &str, name: &str, input_json: &str) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{id}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: id.to_string(), + name: name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn message_text(message: &Message) -> &str { + message + .content + .iter() + .find_map(|block| match block { + ContentBlock::Text { text } => Some(text.as_str()), + _ => None, + }) + .unwrap_or("") +} diff --git a/vendor/mentra/src/agent/tests/runtime_resume.rs b/vendor/mentra/src/agent/tests/runtime_resume.rs new file mode 100644 index 0000000..39aef46 --- /dev/null +++ b/vendor/mentra/src/agent/tests/runtime_resume.rs @@ -0,0 +1,1148 @@ +use std::{ + collections::BTreeMap, + fs, + sync::{ + Arc, + atomic::{AtomicU64, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use serde_json::{Value, json}; +use tokio::sync::{Notify, watch}; + +use crate::{ + BuiltinProvider, ContentBlock, Message, ProviderId, ReasoningFormat, ReasoningProvenance, Role, + TranscriptKind, + agent::{AgentConfig, AgentSnapshot, AgentStatus, TeamConfig}, + provider::{ContentBlockDelta, ContentBlockStart, ProviderEvent}, + runtime::{AgentStore, Runtime, SqliteRuntimeStore}, + team::{TeamMemberStatus, TeamMessage, TeamStore}, + tool::{ + ToolContext, ToolDefinition, ToolDurability, ToolExecutor, ToolOutput, ToolResult, + ToolSideEffectLevel, ToolSpec, + }, +}; + +use super::support::{ScriptedProvider, controlled_stream, erroring_stream, model_info, ok_stream}; + +#[tokio::test] +async fn runtime_startup_preserves_memory_until_resume_and_resume_rolls_back_pending_turn() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("pending-recovery"); + let (script, tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![script], + ); + + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let agent_id = agent.id().to_string(); + let mut snapshot = agent.watch_snapshot(); + + let send_task = tokio::spawn(async move { + let mut agent = agent; + let result = agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await; + (agent, result) + }); + + wait_for_status(&mut snapshot, AgentStatus::Streaming).await; + tx.send(Ok(ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: model.id.clone(), + role: Role::Assistant, + })) + .expect("message started"); + snapshot.changed().await.expect("snapshot changed"); + + tx.send(Ok(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + })) + .expect("block started"); + snapshot.changed().await.expect("snapshot changed"); + + tx.send(Ok(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("Hel".to_string()), + })) + .expect("delta"); + snapshot.changed().await.expect("snapshot changed"); + assert_eq!(snapshot.borrow().current_text, "Hel"); + + send_task.abort(); + let _ = send_task.await; + clear_leases(&store); + + let reboot_runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + Vec::new(), + )) + .build() + .expect("rebuild runtime"); + + let persisted_before_resume = store + .load_agent(&agent_id) + .expect("load interrupted state") + .expect("agent state"); + assert_eq!( + persisted_before_resume + .memory + .pending_turn + .as_ref() + .expect("pending turn persisted") + .current_text, + "Hel" + ); + + let resumed = reboot_runtime + .resume_agent(&agent_id) + .expect("resume interrupted agent"); + assert!(resumed.history().is_empty()); + assert_eq!( + resumed.watch_snapshot().borrow().status, + AgentStatus::Interrupted + ); + assert!(resumed.watch_snapshot().borrow().current_text.is_empty()); + + let persisted_after_resume = store + .load_agent(&agent_id) + .expect("load recovered state") + .expect("agent state"); + assert!(persisted_after_resume.memory.pending_turn.is_none()); + assert!(persisted_after_resume.memory.run.is_none()); + assert!(persisted_after_resume.memory.transcript.is_empty()); +} + +#[tokio::test] +async fn aborted_partial_thinking_stream_rolls_back_instead_of_persisting_empty_signature() { + let provider_id = ProviderId::new("anthropic-edge"); + let model = model_info("claude-requested", provider_id.clone()); + let store = temp_store("aborted-thinking"); + let provider = ScriptedProvider::new( + provider_id.clone(), + vec![model.clone()], + vec![erroring_stream( + vec![ + ProviderEvent::MessageStarted { + id: "msg-thinking".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Thinking { + encrypted_content: None, + id: None, + provenance: Some(ReasoningProvenance { + provider: provider_id, + model: model.id.clone(), + format: ReasoningFormat::AnthropicSigned, + }), + redacted: false, + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingText("partial chain".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ], + crate::ProviderError::MalformedStream("aborted".to_string()), + )], + ); + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let agent_id = agent.id().to_string(); + + agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect_err("aborted stream should fail"); + + assert!(agent.history().is_empty()); + let persisted = store + .load_agent(&agent_id) + .expect("load agent") + .expect("persisted agent"); + assert!(persisted.memory.transcript.is_empty()); +} + +#[tokio::test] +async fn resume_agent_keeps_committed_transcript_when_tool_execution_was_interrupted() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("committed-recovery"); + let tool_started = Arc::new(Notify::new()); + let tool_release = Arc::new(Notify::new()); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![tool_use_stream( + &model.id, + "tool-1", + "blocking_tool", + r#"{"value":"hi"}"#, + )], + ); + + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .with_tool(BlockingTool { + started: tool_started.clone(), + release: tool_release, + }) + .build() + .expect("build runtime"); + let agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let agent_id = agent.id().to_string(); + + let send_task = tokio::spawn(async move { + let mut agent = agent; + let result = agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await; + (agent, result) + }); + + tool_started.notified().await; + send_task.abort(); + let _ = send_task.await; + clear_leases(&store); + + let persisted_before_resume = store + .load_agent(&agent_id) + .expect("load interrupted state") + .expect("agent state"); + assert!(persisted_before_resume.memory.pending_turn.is_none()); + assert!( + persisted_before_resume + .memory + .run + .as_ref() + .expect("run metadata") + .assistant_committed + ); + + let reboot_runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + Vec::new(), + )) + .build() + .expect("rebuild runtime"); + let resumed = reboot_runtime + .resume_agent(&agent_id) + .expect("resume interrupted agent"); + + assert_eq!( + resumed.watch_snapshot().borrow().status, + AgentStatus::Interrupted + ); + assert_eq!(resumed.history().len(), 2); + assert_eq!( + resumed.history()[0], + Message::user(ContentBlock::text("hello")) + ); + assert!(matches!( + &resumed.history()[1].content[0], + ContentBlock::ToolUse { name, .. } if name == "blocking_tool" + )); + + let persisted_after_resume = store + .load_agent(&agent_id) + .expect("load recovered state") + .expect("agent state"); + assert!(persisted_after_resume.memory.run.is_none()); + assert_eq!(persisted_after_resume.memory.transcript.len(), 2); +} + +#[tokio::test] +async fn resume_all_rebuilds_agents_from_agent_memory() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("resume-all"); + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "done")], + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let agent_id = agent.id().to_string(); + + agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await + .expect("send"); + clear_leases(&store); + + let reboot_runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model], + Vec::new(), + )) + .build() + .expect("rebuild runtime"); + let resumed = reboot_runtime.resume_all().expect("resume all"); + + assert_eq!(resumed.len(), 1); + assert_eq!(resumed[0].id(), agent_id); + assert_eq!(resumed[0].history().len(), 2); + assert_eq!( + resumed[0].last_message(), + Some(&Message::assistant(ContentBlock::text("done"))) + ); +} + +#[tokio::test] +async fn signed_and_redacted_thinking_survive_commit_persist_resume_and_replay() { + let provider_id = ProviderId::new("anthropic-edge"); + let model = model_info("claude-requested", provider_id.clone()); + let store = temp_store("thinking-replay"); + let expected_thinking = vec![ + ContentBlock::Thinking { + thinking: "private chain".to_string(), + signature: Some("opaque-signature".to_string()), + encrypted_content: None, + id: None, + provenance: Some(ReasoningProvenance { + provider: provider_id.clone(), + model: model.id.clone(), + format: ReasoningFormat::AnthropicSigned, + }), + redacted: false, + }, + ContentBlock::Thinking { + thinking: String::new(), + signature: Some("opaque-redacted-data".to_string()), + encrypted_content: None, + id: None, + provenance: Some(ReasoningProvenance { + provider: provider_id.clone(), + model: model.id.clone(), + format: ReasoningFormat::AnthropicSigned, + }), + redacted: true, + }, + ]; + let first_provider = ScriptedProvider::new( + provider_id.clone(), + vec![model.clone()], + vec![thinking_stream( + &provider_id, + &model.id, + "private chain", + "opaque-signature", + "opaque-redacted-data", + "visible answer", + )], + ); + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(first_provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let agent_id = agent.id().to_string(); + + let response = agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect("commit thinking response"); + assert_eq!(&response.content[..2], expected_thinking.as_slice()); + assert_eq!(response.text(), "visible answer"); + + let persisted = store + .load_agent(&agent_id) + .expect("load agent") + .expect("persisted agent"); + let persisted_assistant = persisted + .memory + .transcript + .to_messages() + .into_iter() + .find(|message| message.role == Role::Assistant) + .expect("persisted assistant message"); + assert_eq!( + &persisted_assistant.content[..2], + expected_thinking.as_slice() + ); + clear_leases(&store); + + let replay_provider = ScriptedProvider::new( + provider_id, + vec![model.clone()], + vec![text_stream(&model.id, "continued")], + ); + let reboot_runtime = Runtime::empty_builder() + .with_store(store) + .with_provider_instance(replay_provider.clone()) + .build() + .expect("rebuild runtime"); + let mut resumed = reboot_runtime + .resume_agent(&agent_id) + .expect("resume persisted agent"); + + resumed + .send(vec![ContentBlock::text("continue")]) + .await + .expect("replay persisted thinking"); + let requests = replay_provider.recorded_requests().await; + let replayed_assistant = requests[0] + .messages + .iter() + .find(|message| message.role == Role::Assistant) + .expect("replayed assistant message"); + assert_eq!( + &replayed_assistant.content[..2], + expected_thinking.as_slice() + ); +} + +#[tokio::test] +async fn responses_reasoning_and_paired_tool_ids_survive_agent_replay() { + let provider_id = ProviderId::new("openai-edge"); + let model = model_info("gpt-requested", provider_id.clone()); + let provider = ScriptedProvider::new( + provider_id.clone(), + vec![model.clone()], + vec![ + responses_reasoning_tool_stream(&provider_id, &model.id), + text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_store(temp_store("responses-reasoning-replay")) + .with_provider_instance(provider.clone()) + .with_tool(DetailsTool) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("use the tool")]) + .await + .expect("Responses reasoning tool loop should complete"); + + let requests = provider.recorded_requests().await; + assert_eq!(requests.len(), 2); + let replayed_assistant = requests[1] + .messages + .iter() + .find(|message| message.role == Role::Assistant) + .expect("assistant reasoning turn should replay"); + assert!(matches!( + &replayed_assistant.content[0], + ContentBlock::Thinking { + id: Some(id), + encrypted_content: Some(encrypted_content), + .. + } if id == "rs_1" && encrypted_content == "encrypted-1" + )); + assert!(matches!( + &replayed_assistant.content[1], + ContentBlock::ToolUse { id, .. } if id == "call_1|fc_1" + )); + assert!(requests[1].messages.iter().any(|message| { + message.content.iter().any(|block| { + matches!( + block, + ContentBlock::ToolResult { tool_use_id, .. } if tool_use_id == "call_1|fc_1" + ) + }) + })); +} + +// M3 test 1: `ToolOutput::details` survives a real restart — persisted to +// the SQLite `agent_memory` row by a live tool round, then recovered by a +// brand new `AgentMemory` built from `RuntimeStore::load_agent` (the same +// `resume_all` path every crash-recovery test in this file exercises), +// keyed by the originating `tool_use_id`. +#[tokio::test] +async fn resumed_agent_keeps_tool_result_details_after_restart() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("resume-details"); + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "details_tool", r#"{}"#), + text_stream(&model.id, "done"), + ], + )) + .with_tool(DetailsTool) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let agent_id = agent.id().to_string(); + + agent + .send(vec![ContentBlock::Text { + text: "run the details tool".to_string(), + }]) + .await + .expect("send"); + clear_leases(&store); + + let reboot_runtime = Runtime::empty_builder() + .with_store(store) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model], + Vec::new(), + )) + .build() + .expect("rebuild runtime"); + let resumed = reboot_runtime.resume_all().expect("resume all"); + + assert_eq!(resumed.len(), 1); + assert_eq!(resumed[0].id(), agent_id); + + let item = resumed[0] + .transcript() + .items() + .iter() + .find(|item| matches!(item.kind, TranscriptKind::ToolExchange { .. })) + .expect("resumed transcript keeps the tool exchange item"); + assert_eq!( + item.details(), + Some(&BTreeMap::from([( + "call-1".to_string(), + json!({ "marker": "keep-me" }), + )])) + ); +} + +#[tokio::test] +async fn resume_filters_agents_by_runtime_identifier() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("resume-filter"); + + let runtime_a = Runtime::empty_builder() + .with_runtime_identifier("session-a") + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + Vec::new(), + )) + .build() + .expect("build runtime a"); + let agent_a = runtime_a + .spawn("agent-a", model.clone()) + .expect("spawn agent a"); + + let runtime_b = Runtime::empty_builder() + .with_runtime_identifier("session-b") + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + Vec::new(), + )) + .build() + .expect("build runtime b"); + let _agent_b = runtime_b + .spawn("agent-b", model.clone()) + .expect("spawn agent b"); + + clear_leases(&store); + + let reboot_runtime = Runtime::empty_builder() + .with_runtime_identifier("session-a") + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model], + Vec::new(), + )) + .build() + .expect("rebuild runtime"); + + let resumed = reboot_runtime + .resume("session-a") + .expect("resume session-a"); + assert_eq!(resumed.len(), 1); + assert_eq!(resumed[0].id(), agent_a.id()); + assert_eq!(resumed[0].name(), "agent-a"); +} + +#[tokio::test] +async fn list_persisted_agents_includes_teammates_for_runtime() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("persisted-agent-list"); + let runtime_identifier = "persisted-agent-list"; + + let runtime = Runtime::builder() + .with_runtime_identifier(runtime_identifier) + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + Vec::new(), + )) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("persisted-agent-list-team")), + ..Default::default() + }, + ) + .expect("spawn lead"); + lead.spawn_teammate("alice", "researcher", None) + .await + .expect("spawn teammate"); + + let persisted = runtime + .list_persisted_agents(runtime_identifier) + .expect("list persisted agents"); + assert_eq!(persisted.len(), 2); + assert_eq!(persisted[0].name, "lead"); + assert!(!persisted[0].is_teammate); + assert_eq!(persisted[1].name, "alice"); + assert!(persisted[1].is_teammate); +} + +#[tokio::test] +async fn dropping_runtime_releases_agent_lease_for_next_resume() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("lease-release"); + + let runtime = Runtime::empty_builder() + .with_runtime_identifier("lease-release") + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "done")], + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let agent_id = agent.id().to_string(); + + agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await + .expect("send"); + drop(agent); + drop(runtime); + + let reboot_runtime = Runtime::empty_builder() + .with_runtime_identifier("lease-release") + .with_store(store) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model], + Vec::new(), + )) + .build() + .expect("rebuild runtime"); + + let resumed = reboot_runtime + .resume("lease-release") + .expect("resume runtime after drop"); + assert_eq!(resumed.len(), 1); + assert_eq!(resumed[0].id(), agent_id); +} + +#[tokio::test] +async fn resume_revives_persisted_teammate_actors_for_lead_runtime() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("teammate-revive-resume"); + let runtime_identifier = "teammate-revive"; + + let initial_runtime = Runtime::builder() + .with_runtime_identifier(runtime_identifier) + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + Vec::new(), + )) + .build() + .expect("build initial runtime"); + let mut lead = initial_runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(temp_team_dir("resume-revive-team")), + ..Default::default() + }, + ) + .expect("spawn lead"); + lead.spawn_teammate("alice", "researcher", None) + .await + .expect("spawn teammate"); + drop(lead); + drop(initial_runtime); + clear_leases(&store); + + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "revived-send", + "team_send", + r#"{"to":"lead","content":"revived and responsive"}"#, + ), + text_stream(&model.id, "done"), + text_stream(&model.id, "checked"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::builder() + .with_runtime_identifier(runtime_identifier) + .with_store(store) + .with_provider_instance(provider) + .build() + .expect("build resumed runtime"); + + let mut resumed = runtime.resume(runtime_identifier).expect("resume runtime"); + assert_eq!(resumed.len(), 1); + let mut lead = resumed.pop().expect("lead agent"); + assert!(!lead.is_teammate()); + + lead.send_team_message("alice", "Ping me after restart") + .expect("send team message"); + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + lead.send(vec![ContentBlock::Text { + text: "status?".to_string(), + }]) + .await + .expect("send status check"); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 3); + let inbox = latest_team_inbox_text(&requests[2]).expect("team inbox"); + assert!(inbox.contains("alice")); + assert!(inbox.contains("revived and responsive")); +} + +#[tokio::test] +async fn resume_wakes_revived_teammate_for_persisted_inbox_work() { + let model = model_info("model", BuiltinProvider::Anthropic); + let store = temp_store("teammate-revive-pending-inbox"); + let runtime_identifier = "teammate-revive-pending-inbox"; + let team_dir = temp_team_dir("resume-pending-inbox-team"); + + let initial_runtime = Runtime::builder() + .with_runtime_identifier(runtime_identifier) + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + Vec::new(), + )) + .build() + .expect("build initial runtime"); + let mut lead = initial_runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .expect("spawn lead"); + lead.spawn_teammate("alice", "researcher", None) + .await + .expect("spawn teammate"); + drop(lead); + drop(initial_runtime); + clear_leases(&store); + + ::append_team_message( + &store, + team_dir.as_path(), + "alice", + &TeamMessage::message( + "lead".to_string(), + "Handle this persisted inbox work after restart".to_string(), + ), + ) + .expect("append persisted inbox message"); + + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "persisted-inbox-send", + "team_send", + r#"{"to":"lead","content":"processed persisted inbox"}"#, + ), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::builder() + .with_runtime_identifier(runtime_identifier) + .with_store(store) + .with_provider_instance(provider) + .build() + .expect("build resumed runtime"); + + let mut resumed = runtime.resume(runtime_identifier).expect("resume runtime"); + assert_eq!(resumed.len(), 1); + let lead = resumed.pop().expect("lead agent"); + assert!(!lead.is_teammate()); + + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 2); + let inbox = latest_team_inbox_text(&requests[0]).expect("team inbox"); + assert!(inbox.contains("Handle this persisted inbox work after restart")); +} + +/// A tool that returns opaque `ToolOutput::details` metadata, used to prove +/// details survive a real restart (persist/reload through the SQLite store). +struct DetailsTool; + +#[async_trait] +impl ToolDefinition for DetailsTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("details_tool") + .description("test tool: returns opaque details metadata") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .build() + } +} + +#[async_trait] +impl ToolExecutor for DetailsTool { + async fn execute_mut_output( + &self, + _ctx: ToolContext<'_>, + _input: Value, + ) -> Result { + Ok(ToolOutput::text("tool output").with_details(json!({ "marker": "keep-me" }))) + } +} + +struct BlockingTool { + started: Arc, + release: Arc, +} + +#[async_trait] +impl ToolDefinition for BlockingTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("blocking_tool") + .description("blocks until released") + .input_schema(json!({ + "type": "object", + "properties": { + "value": { "type": "string" } + } + })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .build() + } +} + +#[async_trait] +impl ToolExecutor for BlockingTool { + async fn execute_mut(&self, _ctx: ToolContext<'_>, _input: Value) -> ToolResult { + self.started.notify_one(); + self.release.notified().await; + Ok("released".to_string()) + } +} + +async fn wait_for_status(receiver: &mut watch::Receiver, status: AgentStatus) { + loop { + if receiver.borrow().status == status { + return; + } + receiver.changed().await.expect("snapshot changed"); + } +} + +async fn wait_for_teammate_status(agent: &crate::agent::Agent, status: TeamMemberStatus) { + let mut receiver = agent.watch_snapshot(); + loop { + if receiver + .borrow() + .teammates + .iter() + .any(|teammate| teammate.name == "alice" && teammate.status == status) + { + return; + } + receiver.changed().await.expect("snapshot changed"); + } +} + +fn tool_use_stream( + model: &str, + id: &str, + name: &str, + input_json: &str, +) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{id}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: id.to_string(), + name: name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn text_stream(model: &str, text: &str) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn thinking_stream( + provider: &ProviderId, + model: &str, + thinking: &str, + signature: &str, + redacted_data: &str, + text: &str, +) -> super::support::StreamScript { + let provenance = Some(ReasoningProvenance { + provider: provider.clone(), + model: model.to_string(), + format: ReasoningFormat::AnthropicSigned, + }); + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-thinking".to_string(), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Thinking { + encrypted_content: None, + id: None, + provenance: provenance.clone(), + redacted: false, + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingText(thinking.to_string()), + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingSignature(signature.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::ContentBlockStarted { + index: 1, + kind: ContentBlockStart::Thinking { + encrypted_content: None, + id: None, + provenance, + redacted: true, + }, + }, + ProviderEvent::ContentBlockDelta { + index: 1, + delta: ContentBlockDelta::ThinkingSignature(redacted_data.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 1 }, + ProviderEvent::ContentBlockStarted { + index: 2, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 2, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 2 }, + ProviderEvent::MessageStopped, + ]) +} + +fn responses_reasoning_tool_stream( + provider: &ProviderId, + model: &str, +) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "resp-reasoning-tool".to_string(), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Thinking { + encrypted_content: None, + id: Some("rs_1".to_string()), + provenance: Some(ReasoningProvenance { + provider: provider.clone(), + model: model.to_string(), + format: ReasoningFormat::OpenAiEncrypted, + }), + redacted: false, + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingText("short summary".to_string()), + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingEncryptedContent("encrypted-1".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::ContentBlockStarted { + index: 1, + kind: ContentBlockStart::ToolUse { + id: "call_1|fc_1".to_string(), + name: "details_tool".to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 1, + delta: ContentBlockDelta::ToolUseInputJson("{}".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 1 }, + ProviderEvent::MessageStopped, + ]) +} + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +fn temp_store(label: &str) -> SqliteRuntimeStore { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = std::env::temp_dir().join(format!( + "mentra-runtime-resume-{label}-{timestamp}-{unique}.sqlite" + )); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).expect("create temp dir"); + } + SqliteRuntimeStore::new(path) +} + +fn temp_team_dir(label: &str) -> std::path::PathBuf { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = + std::env::temp_dir().join(format!("mentra-runtime-team-{label}-{timestamp}-{unique}")); + fs::create_dir_all(&path).expect("create temp team dir"); + path +} + +fn team_config(team_dir: std::path::PathBuf) -> TeamConfig { + TeamConfig { + team_dir, + ..Default::default() + } +} + +fn clear_leases(store: &SqliteRuntimeStore) { + let conn = rusqlite::Connection::open(store.path()).expect("open store"); + conn.execute("DELETE FROM leases", []) + .expect("clear leases"); +} + +async fn wait_for_recorded_requests(provider: &ScriptedProvider, expected: usize) { + for _ in 0..250 { + if provider.recorded_requests().await.len() >= expected { + return; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + panic!("timed out waiting for {expected} recorded requests"); +} + +fn latest_team_inbox_text(request: &crate::provider::Request<'_>) -> Option { + request.messages.iter().rev().find_map(|message| { + message.content.iter().find_map(|block| match block { + ContentBlock::Text { text } if text.contains("") => Some(text.clone()), + _ => None, + }) + }) +} diff --git a/vendor/mentra/src/agent/tests/runtime_snapshot.rs b/vendor/mentra/src/agent/tests/runtime_snapshot.rs new file mode 100644 index 0000000..3fe20dc --- /dev/null +++ b/vendor/mentra/src/agent/tests/runtime_snapshot.rs @@ -0,0 +1,385 @@ +use std::{ + sync::atomic::{AtomicU64, Ordering}, + time::{SystemTime, UNIX_EPOCH}, +}; + +use tokio::{ + sync::watch, + time::{Duration, timeout}, +}; + +use crate::{ + AgentConfig, BackgroundTaskStatus, BuiltinProvider, ContentBlock, Role, + agent::{AgentSnapshot, AgentStatus, TeamAutonomyConfig, TeamConfig}, + provider::{ContentBlockDelta, ContentBlockStart, ProviderEvent}, + runtime::{Runtime, RuntimePolicy, SqliteRuntimeStore}, +}; + +use super::support::{ + ScriptedProvider, background_success_command, command_input_json, controlled_stream, + model_info, ok_stream, text_stream, +}; + +#[tokio::test] +async fn owned_waits_coexist_with_mutable_runs_and_track_run_generation() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first"), + text_stream(&model.id, "second"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let first_idle = agent.wait_until_idle(); + let first_finished = agent.wait_for_snapshot(|snapshot| { + snapshot.run_generation == 1 && snapshot.status == AgentStatus::Finished + }); + let (first_snapshot, predicate_snapshot, first_result) = tokio::join!( + first_idle, + first_finished, + agent.send(vec![ContentBlock::text("first run")]) + ); + assert_eq!(first_result.expect("first run").text(), "first"); + assert_eq!(first_snapshot.run_generation, 1); + assert_eq!(predicate_snapshot.run_generation, 1); + assert_eq!(first_snapshot.status, AgentStatus::Finished); + + // Constructed while the previous generation is terminal: this must wait + // for generation 2 rather than immediately returning generation 1. + let second_idle = agent.wait_until_idle(); + let (second_snapshot, second_result) = tokio::join!( + second_idle, + agent.send(vec![ContentBlock::text("second run")]) + ); + assert_eq!(second_result.expect("second run").text(), "second"); + assert_eq!(second_snapshot.run_generation, 2); + assert_eq!(second_snapshot.status, AgentStatus::Finished); +} + +#[tokio::test] +async fn wait_handle_targets_the_generation_active_when_the_future_is_created() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (script, tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![script], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let waits = agent.wait_handle(); + let mut snapshots = agent.watch_snapshot(); + + let send_task = tokio::spawn(async move { + let mut agent = agent; + agent.send(vec![ContentBlock::text("start")]).await + }); + wait_for_status(&mut snapshots, AgentStatus::Streaming).await; + assert_eq!(snapshots.borrow().run_generation, 1); + + let idle = waits.wait_until_idle(); + let finished = waits.wait_for_snapshot(|snapshot| { + snapshot.run_generation == 1 && snapshot.status == AgentStatus::Finished + }); + for event in [ + ProviderEvent::MessageStarted { + id: "msg-active-wait".to_string(), + model: model.id, + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("done".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ] { + tx.send(Ok(event)).expect("stream receiver remains alive"); + } + drop(tx); + + let (idle, finished, result) = tokio::join!( + timeout(Duration::from_secs(5), idle), + timeout(Duration::from_secs(5), finished), + send_task, + ); + let idle = idle.expect("idle wait timed out"); + let finished = finished.expect("snapshot wait timed out"); + assert_eq!( + result + .expect("send task joins") + .expect("run succeeds") + .text(), + "done" + ); + assert_eq!(idle.run_generation, 1); + assert_eq!(idle.status, AgentStatus::Finished); + assert_eq!(finished.run_generation, 1); + assert_eq!(finished.status, AgentStatus::Finished); +} + +#[tokio::test] +async fn teammate_reply_wait_consumes_the_snapshot_signaled_inbox() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let team_dir = std::env::temp_dir().join(format!( + "mentra-wait-team-{}-{}", + std::process::id(), + NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed) + )); + let config = AgentConfig { + team: TeamConfig { + team_dir, + autonomy: TeamAutonomyConfig::default(), + }, + ..AgentConfig::default() + }; + let alice = runtime + .spawn_with_config("alice", model.clone(), config.clone()) + .expect("spawn alice"); + let bob = runtime + .spawn_with_config("bob", model, config) + .expect("spawn bob"); + let waits = bob.wait_handle(); + let reply = bob.wait_for_teammate_reply(); + + alice + .send_team_message("bob", "the review is ready") + .expect("send reply"); + let messages = timeout(Duration::from_secs(5), reply) + .await + .expect("reply wait timed out") + .expect("read reply"); + + assert_eq!(messages.len(), 1); + assert_eq!(messages[0].sender, "alice"); + assert_eq!(messages[0].content, "the review is ready"); + assert_eq!(bob.watch_snapshot().borrow().pending_team_messages, 0); + + let reply = waits.wait_for_teammate_reply(); + alice + .send_team_message("bob", "the follow-up review is ready") + .expect("send second reply"); + let messages = timeout(Duration::from_secs(5), reply) + .await + .expect("handle reply wait timed out") + .expect("read second reply"); + + assert_eq!(messages.len(), 1); + assert_eq!(messages[0].sender, "alice"); + assert_eq!(messages[0].content, "the follow-up review is ready"); + assert_eq!(bob.watch_snapshot().borrow().pending_team_messages, 0); +} + +#[tokio::test] +async fn snapshot_progresses_during_streaming() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (script, tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![script], + ); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let agent = runtime.spawn("agent", model.clone()).unwrap(); + let mut snapshot = agent.watch_snapshot(); + + let send_task = tokio::spawn(async move { + let mut agent = agent; + let result = agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await; + (agent, result) + }); + + wait_for_status(&mut snapshot, AgentStatus::Streaming).await; + + tx.send(Ok(ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: model.id, + role: Role::Assistant, + })) + .unwrap(); + snapshot.changed().await.unwrap(); + + tx.send(Ok(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + })) + .unwrap(); + snapshot.changed().await.unwrap(); + + tx.send(Ok(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("Hel".to_string()), + })) + .unwrap(); + snapshot.changed().await.unwrap(); + assert_eq!(snapshot.borrow().current_text, "Hel"); + + tx.send(Ok(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("lo".to_string()), + })) + .unwrap(); + snapshot.changed().await.unwrap(); + assert_eq!(snapshot.borrow().current_text, "Hello"); + + tx.send(Ok(ProviderEvent::ContentBlockStopped { index: 0 })) + .unwrap(); + tx.send(Ok(ProviderEvent::MessageStopped)).unwrap(); + drop(tx); + + let (agent, result) = send_task.await.unwrap(); + result.unwrap(); + + let snapshot = agent.watch_snapshot(); + assert_eq!(snapshot.borrow().status, AgentStatus::Finished); + assert!(snapshot.borrow().current_text.is_empty()); + assert!(snapshot.borrow().pending_tool_uses.is_empty()); +} + +#[tokio::test] +async fn snapshot_updates_when_background_task_finishes() { + let command = background_success_command("bg-done", 50); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-bg".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "tool-bg".to_string(), + name: "background_run".to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(command_input_json(&command)), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-follow".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("continued".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + ], + ); + + let runtime = Runtime::builder() + .with_store(temp_store("snapshot-background-finish")) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + let mut snapshot = agent.watch_snapshot(); + + agent + .send(vec![ContentBlock::Text { + text: "run background command".to_string(), + }]) + .await + .unwrap(); + + wait_for_background_status(&mut snapshot, BackgroundTaskStatus::Finished).await; + assert_eq!(snapshot.borrow().background_tasks.len(), 1); + assert!( + snapshot.borrow().background_tasks[0] + .output_preview + .as_deref() + .is_some_and(|preview| preview.contains("bg-done")) + ); +} + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +fn temp_store(label: &str) -> SqliteRuntimeStore { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + SqliteRuntimeStore::new(std::env::temp_dir().join(format!( + "mentra-runtime-store-{label}-{timestamp}-{unique}.sqlite" + ))) +} + +async fn wait_for_status(receiver: &mut watch::Receiver, status: AgentStatus) { + timeout(Duration::from_secs(90), async { + loop { + if receiver.borrow().status == status { + return; + } + receiver.changed().await.unwrap(); + } + }) + .await + .unwrap_or_else(|_| panic!("timed out waiting for agent status {status:?}")); +} + +async fn wait_for_background_status( + receiver: &mut watch::Receiver, + status: BackgroundTaskStatus, +) { + timeout(Duration::from_secs(90), async { + loop { + if receiver + .borrow() + .background_tasks + .iter() + .any(|task| task.status == status) + { + return; + } + receiver.changed().await.unwrap(); + } + }) + .await + .unwrap_or_else(|_| panic!("timed out waiting for background status {status:?}")); +} diff --git a/vendor/mentra/src/agent/tests/runtime_tasks.rs b/vendor/mentra/src/agent/tests/runtime_tasks.rs new file mode 100644 index 0000000..cc148f1 --- /dev/null +++ b/vendor/mentra/src/agent/tests/runtime_tasks.rs @@ -0,0 +1,502 @@ +use std::{ + fs, + path::PathBuf, + sync::atomic::{AtomicU64, Ordering}, + time::{SystemTime, UNIX_EPOCH}, +}; + +use crate::{ + BuiltinProvider, ContentBlock, Message, Role, + agent::{AgentConfig, CompactionConfig, TaskConfig}, + provider::{ContentBlockDelta, ContentBlockStart, ProviderError, ProviderEvent}, + runtime::{ + NewTask, Runtime, SqliteRuntimeStore, TaskItem, TaskPatch, TaskStatus, TaskStore, + task::TASK_REMINDER_TEXT, + }, +}; + +use super::super::TeammateIdentity; +use super::support::{ScriptedProvider, erroring_stream, model_info, ok_stream}; + +#[test] +fn task_board_exposes_typed_lead_operations_and_agent_namespace() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_tasks_dir("typed-board"); + let runtime = Runtime::builder() + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![], + )) + .build() + .expect("build runtime"); + let board = runtime.task_board(&tasks_dir); + + let first = board + .create(NewTask { + owner: "alice".to_string(), + ..NewTask::new("design") + }) + .expect("create first task"); + let second = board + .create(NewTask { + working_directory: Some("workspace".to_string()), + ..NewTask::new("implement") + }) + .expect("create second task"); + assert_eq!((first.id, second.id), (1, 2)); + + let dependent = board + .add_dependency(first.id, second.id) + .expect("add dependency"); + assert_eq!(dependent.blocked_by, vec![first.id]); + assert_eq!( + board.get(first.id).expect("get first").blocks, + vec![second.id] + ); + assert_eq!(board.list().expect("list tasks").len(), 2); + + let unblocked = board + .remove_dependency(first.id, second.id) + .expect("remove dependency"); + assert!(unblocked.blocked_by.is_empty()); + let updated = board + .update( + second.id, + TaskPatch { + status: Some(TaskStatus::InProgress), + working_directory: Some(None), + ..TaskPatch::default() + }, + ) + .expect("update task"); + assert_eq!(updated.status, TaskStatus::InProgress); + assert_eq!(updated.working_directory, None); + + let deserialized_patch: TaskPatch = serde_json::from_value(serde_json::json!({ + "workingDirectory": null + })) + .expect("deserialize explicit task working-directory clear"); + assert_eq!(deserialized_patch.working_directory, Some(None)); + + let agent = runtime + .spawn_with_config("lead", model, task_config(tasks_dir)) + .expect("spawn agent"); + assert_eq!( + agent.task_board().list().expect("agent board list").len(), + 2 + ); +} + +#[test] +fn task_board_preserves_teammate_access_and_explicit_lead_claimant() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_tasks_dir("typed-board-access"); + let runtime = Runtime::builder() + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![], + )) + .build() + .expect("build runtime"); + let lead_board = runtime.task_board(&tasks_dir); + let alice_task = lead_board + .create(NewTask { + owner: "alice".to_string(), + ..NewTask::new("alice-owned") + }) + .expect("create alice task"); + let bob_task = lead_board + .create(NewTask { + owner: "bob".to_string(), + ..NewTask::new("bob-owned") + }) + .expect("create bob task"); + let unowned = lead_board + .create(NewTask::new("claimable")) + .expect("create claimable task"); + + let mut alice = runtime + .spawn_with_config("alice", model, task_config(tasks_dir)) + .expect("spawn alice"); + alice.teammate_identity = Some(TeammateIdentity { + role: "worker".to_string(), + lead: "lead".to_string(), + }); + let alice_board = alice.task_board(); + + alice_board + .update( + alice_task.id, + TaskPatch { + description: Some("mine".to_string()), + ..TaskPatch::default() + }, + ) + .expect("update own task"); + assert!( + alice_board + .update( + bob_task.id, + TaskPatch { + description: Some("not mine".to_string()), + ..TaskPatch::default() + }, + ) + .expect_err("cannot update another teammate task") + .to_string() + .contains("cannot update") + ); + assert!( + alice_board + .add_dependency(alice_task.id, bob_task.id) + .expect_err("teammate cannot edit dependencies") + .to_string() + .contains("cannot update") + ); + assert!( + alice_board + .claim(Some(unowned.id), "bob") + .expect_err("teammate cannot pose as bob") + .to_string() + .contains("cannot claim a task for") + ); + + let claimed = lead_board + .claim(Some(unowned.id), "carol") + .expect("lead board uses explicit claimant"); + assert_eq!(claimed.owner, "carol"); +} + +#[tokio::test] +async fn task_updates_snapshot_and_persists_for_new_agents() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_tasks_dir("persist"); + let store = temp_store("persist"); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + task_tool_stream( + "tool-1", + "task_create", + r#"{"subject":"Plan work","owner":"agent-a"}"#, + ), + text_stream("created"), + ], + ); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let config = task_config(tasks_dir.clone()); + let mut agent = runtime + .spawn_with_config("agent", model.clone(), config.clone()) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "start".to_string(), + }]) + .await + .expect("send"); + + assert_eq!( + agent.watch_snapshot().borrow().tasks, + vec![TaskItem { + id: 1, + subject: "Plan work".to_string(), + description: String::new(), + status: TaskStatus::Pending, + blocked_by: Vec::new(), + blocks: Vec::new(), + owner: "agent-a".to_string(), + working_directory: None, + }] + ); + assert_eq!( + store + .load_tasks(tasks_dir.as_path()) + .expect("load persisted tasks") + .len(), + 1 + ); + + let other_provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream("ok")], + ); + let other_runtime = Runtime::builder() + .with_store(store) + .with_provider_instance(other_provider) + .build() + .expect("build runtime"); + let other_agent = other_runtime + .spawn_with_config("other", model, config) + .expect("spawn other agent"); + assert_eq!(other_agent.watch_snapshot().borrow().tasks.len(), 1); +} + +#[tokio::test] +async fn task_reminder_is_injected_after_three_rounds_without_task_tools() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_tasks_dir("reminder"); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + task_tool_stream("tool-1", "task_create", r#"{"subject":"Plan work"}"#), + text_stream("created"), + text_stream("round 1"), + text_stream("round 2"), + text_stream("round 3"), + text_stream("round 4"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + system: Some("Base system prompt".to_string()), + task: TaskConfig { + tasks_dir, + reminder_threshold: 3, + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "set task".to_string(), + }]) + .await + .expect("create task"); + + for round in 1..=4 { + agent + .send(vec![ContentBlock::Text { + text: format!("round {round}"), + }]) + .await + .expect("send round"); + } + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 6); + assert_eq!(requests[0].system.as_deref(), Some("Base system prompt")); + assert_eq!(requests[3].system.as_deref(), Some("Base system prompt")); + + let expected_system = format!("{TASK_REMINDER_TEXT}\n\nBase system prompt"); + assert_eq!( + requests[4].system.as_deref(), + Some(expected_system.as_str()) + ); + assert_eq!( + requests[5].system.as_deref(), + Some(expected_system.as_str()) + ); +} + +#[tokio::test] +async fn task_state_rolls_back_when_run_fails() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_tasks_dir("rollback"); + let store = temp_store("rollback"); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + task_tool_stream("tool-1", "task_create", r#"{"subject":"Plan work"}"#), + erroring_stream( + vec![ProviderEvent::MessageStarted { + id: "msg-fail".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }], + ProviderError::MalformedStream("boom".to_string()), + ), + ], + ); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("agent", model, task_config(tasks_dir.clone())) + .expect("spawn agent"); + + let result = agent + .send(vec![ContentBlock::Text { + text: "create task".to_string(), + }]) + .await; + + assert!(result.is_err()); + assert!(agent.history().is_empty()); + assert!(agent.watch_snapshot().borrow().tasks.is_empty()); + assert!( + store + .load_tasks(tasks_dir.as_path()) + .expect("load rolled-back tasks") + .is_empty() + ); +} + +#[tokio::test] +async fn task_survives_auto_compaction() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_tasks_dir("compact"); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + task_tool_stream("tool-1", "task_create", r#"{"subject":"Plan work"}"#), + text_stream("created"), + text_stream("summary"), + text_stream("after compact"), + ], + ); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + task: TaskConfig { + tasks_dir, + reminder_threshold: 3, + }, + compaction: CompactionConfig { + auto_compact_threshold_tokens: Some(500), + ..CompactionConfig::default() + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "create task".to_string(), + }]) + .await + .expect("create task"); + agent + .send(vec![ContentBlock::Text { + text: "trigger compact ".repeat(100), + }]) + .await + .expect("trigger compact"); + + assert_eq!(agent.watch_snapshot().borrow().tasks.len(), 1); + assert!(agent.history().iter().any(|message| { + matches!( + message, + Message { + role: Role::User, + content, + } if matches!(content.first(), Some(ContentBlock::Text { text }) if text.contains("[Compaction summary]")) + ) + })); +} + +fn task_config(tasks_dir: PathBuf) -> AgentConfig { + AgentConfig { + task: TaskConfig { + tasks_dir, + reminder_threshold: 3, + }, + ..Default::default() + } +} + +fn task_tool_stream( + tool_id: &str, + tool_name: &str, + input_json: &str, +) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{tool_id}"), + model: "model".to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: tool_id.to_string(), + name: tool_name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn text_stream(text: &str) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: "model".to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +fn temp_tasks_dir(label: &str) -> PathBuf { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = + std::env::temp_dir().join(format!("mentra-task-runtime-{label}-{timestamp}-{unique}")); + fs::create_dir_all(&path).expect("create temp dir"); + path +} + +fn temp_store(label: &str) -> SqliteRuntimeStore { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + SqliteRuntimeStore::new(std::env::temp_dir().join(format!( + "mentra-task-runtime-store-{label}-{timestamp}-{unique}.sqlite" + ))) +} diff --git a/vendor/mentra/src/agent/tests/runtime_tools.rs b/vendor/mentra/src/agent/tests/runtime_tools.rs new file mode 100644 index 0000000..f26d97f --- /dev/null +++ b/vendor/mentra/src/agent/tests/runtime_tools.rs @@ -0,0 +1,5441 @@ +use async_trait::async_trait; +use serde_json::json; +use std::{ + fs, + path::{Path, PathBuf}, + sync::{ + Arc, + atomic::{AtomicU64, AtomicUsize, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; +#[cfg(unix)] +use tokio::time::timeout; +use tokio::time::{Duration, sleep}; + +use crate::{ + BackgroundTaskStatus, BuiltinProvider, ContentBlock, FileToolProfile, Message, Role, + agent::{ + Agent, AgentConfig, AgentEvent, MemoryConfig, SpawnedAgentStatus, TaskConfig, + TeamAutonomyConfig, TeamConfig, ToolProfile, WorkspaceConfig, + }, + memory::{MemoryRecord, MemoryRecordKind, MemoryStore}, + provider::{ + ContentBlockDelta, ContentBlockStart, ProviderError, ProviderEvent, Request, ToolChoice, + ToolSearchMode, + }, + runtime::{ + CancellationToken, HybridRuntimeStore, RunOptions, Runtime, RuntimeError, RuntimePolicy, + SqliteRuntimeStore, TaskIntrinsicTool, + control::{HookDecision, PreExecutionContext, PreExecutionHook}, + task::{self, TaskAccess}, + }, + team::{TeamMemberStatus, TeamMessageKind, TeamProtocolStatus}, +}; + +use super::support::{ + ProbeTool, ScriptedProvider, StaticTool, StopTrippingTool, StreamScript, + background_failure_command, background_success_command, command_input_json, + command_input_with_working_directory_json, controlled_stream, erroring_stream, model_info, + ok_stream, shell_pwd_command, +}; + +#[tokio::test] +async fn send_tool_use_turn_executes_tool_and_commits_follow_up_response() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "tool-1".to_string(), + name: "echo_tool".to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(r#"{"value":"hi"}"#.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-2".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("done".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + ], + ); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("echo_tool", "tool output")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "hi".to_string(), + }]) + .await + .unwrap(); + + assert_eq!(agent.history().len(), 4); + assert_eq!( + agent.history()[2], + Message::user(ContentBlock::ToolResult { + tool_use_id: "tool-1".to_string(), + content: "tool output".into(), + is_error: false, + }) + ); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("done"))) + ); + + let events = collect_events(&mut events); + assert!( + events + .iter() + .any(|event| matches!(event, AgentEvent::ToolUseReady { .. })) + ); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::ToolExecutionFinished { + result: ContentBlock::ToolResult { + is_error: false, + .. + } + } + ))); +} + +#[tokio::test] +async fn tool_execution_error_is_wrapped_and_loop_continues() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "tool-1".to_string(), + name: "failing_tool".to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(r#"{"value":"hi"}"#.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-2".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("handled".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + ], + ); + + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::failure("failing_tool", "tool failed")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hi".to_string(), + }]) + .await + .unwrap(); + + assert_eq!( + agent.history()[2], + Message::user(ContentBlock::ToolResult { + tool_use_id: "tool-1".to_string(), + content: "tool failed".into(), + is_error: true, + }) + ); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("handled"))) + ); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 2); + assert!(matches!( + &requests[1].messages[2].content[0], + ContentBlock::ToolResult { + tool_use_id, + content, + is_error: true, + } if tool_use_id == "tool-1" && content.to_display_string() == "tool failed" + )); +} + +#[tokio::test] +async fn malformed_tool_json_is_reported_back_to_model_instead_of_aborting() { + let model = model_info("model", BuiltinProvider::OpenAI); + let provider = ScriptedProvider::new( + BuiltinProvider::OpenAI, + vec![model.clone()], + vec![ + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "tool-1".to_string(), + name: "files".to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson( + r#"{"path":"src/main.rs"#.to_string(), + ), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + text_stream(&model.id, "recovered"), + ], + ); + + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "inspect the file".to_string(), + }]) + .await + .unwrap(); + + assert_eq!(agent.history().len(), 3); + assert_eq!( + agent.history()[1], + Message::user(ContentBlock::text( + "One or more tool calls could not be executed because their JSON arguments were invalid. Please retry with valid JSON that matches the tool schema exactly.\n\nTool 'files' (tool-1) failed to parse: EOF while parsing a string at line 1 column 20.\nRaw arguments (truncated): {\"path\":\"src/main.rs" + )) + ); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("recovered"))) + ); + + let events = collect_events(&mut events); + assert!( + !events + .iter() + .any(|event| matches!(event, AgentEvent::ToolUseReady { .. })) + ); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 2); + assert_eq!(requests[1].messages.len(), 2); + assert_eq!( + requests[1].messages[1], + Message::user(ContentBlock::text( + "One or more tool calls could not be executed because their JSON arguments were invalid. Please retry with valid JSON that matches the tool schema exactly.\n\nTool 'files' (tool-1) failed to parse: EOF while parsing a string at line 1 column 20.\nRaw arguments (truncated): {\"path\":\"src/main.rs" + )) + ); +} + +#[tokio::test] +async fn background_run_tool_starts_task_and_continues_the_turn() { + let command = background_success_command("bg-done", 200); + let input = command_input_json(&command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-bg", "background_run", &input), + text_stream(&model.id, "continued"), + ], + ); + + let runtime = Runtime::builder() + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "run background command".to_string(), + }]) + .await + .unwrap(); + let cwd = agent.config().workspace.base_dir.display().to_string(); + + assert_eq!( + agent.history()[2], + Message::user(ContentBlock::ToolResult { + tool_use_id: "tool-bg".to_string(), + content: format!("Started background task bg-1 in {cwd} for `{command}`").into(), + is_error: false, + }) + ); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("continued"))) + ); + + let background_tasks = agent.watch_snapshot().borrow().background_tasks.clone(); + assert_eq!(background_tasks.len(), 1); + assert_eq!(background_tasks[0].status, BackgroundTaskStatus::Running); + + let events = collect_events(&mut events); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::BackgroundTaskStarted { task } + if task.id == "bg-1" && task.command == command + ))); +} + +#[tokio::test] +async fn completed_background_results_are_injected_on_next_send() { + let command = background_success_command("bg-done", 50); + let input = command_input_json(&command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-bg", "background_run", &input), + text_stream(&model.id, "continued"), + text_stream(&model.id, "next turn"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_store(temp_store("bg-results-next-send")) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "run background command".to_string(), + }]) + .await + .unwrap(); + wait_for_background_tasks(&agent, 1, BackgroundTaskStatus::Finished).await; + + agent + .send(vec![ContentBlock::Text { + text: "what finished?".to_string(), + }]) + .await + .unwrap(); + + let requests = provider_handle.recorded_requests().await; + let injected = latest_background_results_text(&requests[2]).expect("background results"); + assert!(injected.contains("")); + assert!(injected.contains("[bg:bg-1] status=finished")); + assert!(injected.contains(&format!("command=\"{command}\""))); + assert!(injected.contains("output=\"bg-done")); +} + +#[tokio::test] +async fn teammate_auto_wakes_after_background_task_finishes() { + let command = background_success_command("bg-done", 50); + let input = command_input_json(&command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-bg", "background_run", &input), + text_stream(&model.id, "started"), + text_stream(&model.id, "processed background result"), + ], + ); + let provider_handle = provider.clone(); + + let workspace_dir = temp_team_dir("teammate-background-autowake"); + let workspace_dir = fs::canonicalize(&workspace_dir).expect("canonicalize workspace dir"); + let team_dir = temp_team_dir("teammate-background-autowake-team"); + let store = temp_store("teammate-background-autowake"); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_policy(RuntimePolicy::permissive().with_allowed_working_root(std::env::temp_dir())) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir), + workspace: workspace_config(&workspace_dir), + ..Default::default() + }, + ) + .expect("spawn lead"); + + lead.spawn_teammate( + "alice", + "coder", + Some("Start a background command, then keep going when it finishes.".to_string()), + ) + .await + .expect("spawn teammate"); + + let teammate_id = lead.watch_snapshot().borrow().teammates[0].id.clone(); + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_background_task_record(&store, &teammate_id, 1).await; + wait_for_background_task_status(&store, &teammate_id, "bg-1", BackgroundTaskStatus::Finished) + .await; + wait_for_recorded_requests(&provider_handle, 3).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + let requests = provider_handle.recorded_requests().await; + let injected = latest_background_results_text(&requests[2]).expect("background results"); + assert!(injected.contains("[bg:bg-1] status=finished")); + assert!(request_contains_text( + &requests[2], + "Review any completed background task results" + )); +} + +#[tokio::test] +async fn check_background_reports_single_task_and_lists_all_tasks() { + let command = background_success_command("bg-done", 50); + let input = command_input_json(&command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-bg", "background_run", &input), + text_stream(&model.id, "started"), + multi_tool_use_stream( + &model.id, + &[ + ("check-one", "check_background", r#"{"task_id":"bg-1"}"#), + ("check-all", "check_background", r#"{}"#), + ], + ), + text_stream(&model.id, "checked"), + ], + ); + + let runtime = Runtime::builder() + .with_store(temp_store("bg-check-reports")) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "run background command".to_string(), + }]) + .await + .unwrap(); + wait_for_background_tasks(&agent, 1, BackgroundTaskStatus::Finished).await; + + agent + .send(vec![ContentBlock::Text { + text: "check it".to_string(), + }]) + .await + .unwrap(); + let cwd = agent.config().workspace.base_dir.display().to_string(); + + let message = &agent.history()[7]; + assert_eq!(message.role, Role::User); + assert_eq!(message.content.len(), 2); + match (&message.content[0], &message.content[1]) { + ( + ContentBlock::ToolResult { + tool_use_id: check_one_id, + content: check_one_content, + is_error: false, + }, + ContentBlock::ToolResult { + tool_use_id: check_all_id, + content: check_all_content, + is_error: false, + }, + ) => { + assert_eq!(check_one_id, "check-one"); + assert_eq!(check_all_id, "check-all"); + assert!( + check_one_content.contains(&format!("[finished] cwd={cwd}\n{command}\nbg-done")) + ); + assert!(check_all_content.contains(&format!("bg-1: [finished] cwd={cwd} {command}"))); + } + other => panic!("unexpected tool result payloads: {other:?}"), + } +} + +#[tokio::test] +async fn task_working_directory_routes_shell_for_teammate() { + let command = shell_pwd_command(); + let input = command_input_json(&command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "pwd", "shell", &input), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + + let repo_dir = temp_team_dir("teammate-working-dir"); + let working_dir = repo_dir.join("alice-task"); + fs::create_dir_all(&working_dir).expect("create working dir"); + let team_dir = temp_team_dir("teammate-context-team"); + let tasks_dir = repo_dir.join(".tasks"); + let store = temp_store("teammate-working-dir"); + create_task_with_directory( + &store, + &tasks_dir, + "Implement feature", + "alice", + vec![], + Some("alice-task"), + ); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir), + task: TaskConfig { + tasks_dir: tasks_dir.clone(), + reminder_threshold: 3, + }, + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn lead"); + + lead.spawn_teammate( + "alice", + "coder", + Some("Set up the task and verify cwd.".to_string()), + ) + .await + .expect("spawn teammate"); + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + let requests = provider_handle.recorded_requests().await; + assert!(request_contains_tool_result( + &requests[1], + working_dir.to_string_lossy().as_ref() + )); + assert_eq!( + load_task(&store, &tasks_dir, 1)["workingDirectory"].as_str(), + Some("alice-task") + ); +} + +#[tokio::test] +async fn teammate_shell_without_working_directory_uses_base_dir() { + let command = shell_pwd_command(); + let input = command_input_json(&command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "pwd", "shell", &input), + text_stream(&model.id, "handled"), + ], + ); + let provider_handle = provider.clone(); + + let repo_dir = temp_team_dir("missing-working-dir"); + let team_dir = temp_team_dir("missing-context-team"); + let tasks_dir = repo_dir.join(".tasks"); + let store = temp_store("missing-working-dir"); + create_task(&store, &tasks_dir, "Implement feature", "alice", vec![]); + + let runtime = Runtime::builder() + .with_store(store) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir), + task: TaskConfig { + tasks_dir, + reminder_threshold: 3, + }, + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn lead"); + + lead.spawn_teammate("alice", "coder", Some("Try to run pwd.".to_string())) + .await + .expect("spawn teammate"); + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + let requests = provider_handle.recorded_requests().await; + assert!(request_contains_tool_result( + &requests[1], + repo_dir.to_string_lossy().as_ref() + )); +} + +#[tokio::test] +async fn shell_working_directory_overrides_default_routing() { + let command = shell_pwd_command(); + let input = command_input_with_working_directory_json(&command, "custom"); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "pwd", "shell", &input), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + + let repo_dir = temp_team_dir("explicit-working-dir"); + let working_dir = repo_dir.join("custom"); + fs::create_dir_all(&working_dir).expect("create working dir"); + let runtime = Runtime::builder() + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + task: TaskConfig { + tasks_dir: repo_dir.join(".tasks"), + reminder_threshold: 3, + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "Create a context and inspect it.".to_string(), + }]) + .await + .expect("send"); + + let requests = provider_handle.recorded_requests().await; + assert!(request_contains_tool_result( + &requests[1], + working_dir.to_string_lossy().as_ref() + )); +} + +#[tokio::test] +async fn shell_tool_is_denied_by_default_policy_and_audited() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "pwd", "shell", r#"{"command":"pwd"}"#), + text_stream(&model.id, "done"), + ], + ); + let store = temp_store("policy-audit"); + let store_path = store.path().to_path_buf(); + let runtime = Runtime::builder() + .with_store(store) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "try shell".to_string(), + }]) + .await + .expect("send"); + + let conn = rusqlite::Connection::open(store_path).expect("open audit db"); + let payload: String = conn + .query_row( + "SELECT payload_json FROM audit_events WHERE event_type = 'authorization_denied' ORDER BY created_at DESC, id DESC LIMIT 1", + [], + |row| row.get(0), + ) + .expect("load audit payload"); + assert!(payload.contains("\"action\":\"shell_command\"")); + assert!(payload.contains("\"agent_id\":\"agent")); +} + +#[tokio::test] +async fn files_tool_reads_numbered_lines() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-read", + "files", + r#"{"operations":[{"op":"read","path":"note.txt","offset":2,"limit":1}]}"#, + ), + text_stream(&model.id, "done"), + ], + ); + + let repo_dir = temp_team_dir("files-read"); + fs::write(repo_dir.join("note.txt"), "alpha\nbeta\ngamma\n").expect("write note"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "read the second line".to_string(), + }]) + .await + .expect("send"); + + assert_eq!( + agent.history()[2], + Message::user(ContentBlock::ToolResult { + tool_use_id: "files-read".to_string(), + content: "read note.txt\nL2: beta".into(), + is_error: false, + }) + ); +} + +#[tokio::test] +async fn oversized_files_read_is_truncated_by_runtime_policy() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-read-truncated", + "files", + r#"{"operations":[{"op":"read","path":"note.txt"}]}"#, + ), + text_stream(&model.id, "done"), + ], + ); + + let repo_dir = temp_team_dir("files-read-truncated"); + fs::write(repo_dir.join("note.txt"), "alpha\nbeta\ngamma\n").expect("write note"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_policy( + RuntimePolicy::default() + .with_max_tool_result_bytes(usize::MAX) + .with_max_tool_result_lines(2) + .spill_full_tool_output(false), + ) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "read the note".to_string(), + }]) + .await + .expect("send"); + + let content = match &agent.history()[2] { + Message { + role: Role::User, + content, + } => match content.first().expect("tool result") { + ContentBlock::ToolResult { + content, is_error, .. + } => { + assert!(!is_error); + content.to_display_string() + } + other => panic!("unexpected content block: {other:?}"), + }, + other => panic!("expected tool result message, got {other:?}"), + }; + + fs::remove_dir_all(&repo_dir).expect("remove files-read workspace"); + assert_eq!( + content, + "read note.txt\nL1: alpha\n[truncated: showing 2 of 4 lines; full output was not saved because spill-to-file is disabled by runtime policy]" + ); +} + +#[tokio::test] +async fn files_tool_can_stage_create_then_read_in_one_call() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-create-read", + "files", + r#"{"operations":[{"op":"create","path":"draft.txt","content":"hello"},{"op":"read","path":"draft.txt"}]}"#, + ), + text_stream(&model.id, "done"), + ], + ); + + let repo_dir = temp_team_dir("files-create-read"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "create a draft and read it back".to_string(), + }]) + .await + .expect("send"); + + let content = match &agent.history()[2] { + Message { + role: Role::User, + content, + } => content.first().expect("tool result"), + _ => panic!("expected tool result"), + }; + match content { + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => { + assert_eq!(tool_use_id, "files-create-read"); + assert!(!is_error); + assert!(content.contains("create draft.txt")); + assert!(content.contains("read draft.txt")); + assert!(content.contains("L1: hello")); + } + other => panic!("unexpected content block: {other:?}"), + } + assert_eq!( + fs::read_to_string(repo_dir.join("draft.txt")).expect("read draft"), + "hello" + ); +} + +#[tokio::test] +async fn files_tool_can_update_existing_files() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-set-read", + "files", + r#"{"operations":[{"op":"set","path":"note.txt","content":"updated\n"},{"op":"read","path":"note.txt"}]}"#, + ), + text_stream(&model.id, "done"), + ], + ); + + let repo_dir = temp_team_dir("files-set-read"); + fs::write(repo_dir.join("note.txt"), "original\n").expect("write note"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "update the note and read it back".to_string(), + }]) + .await + .expect("send"); + + let content = match &agent.history()[2] { + Message { + role: Role::User, + content, + } => content.first().expect("tool result"), + _ => panic!("expected tool result"), + }; + match content { + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => { + assert_eq!(tool_use_id, "files-set-read"); + assert!(!is_error); + assert!(content.contains("set note.txt")); + assert!(content.contains("L1: updated")); + } + other => panic!("unexpected content block: {other:?}"), + } + assert_eq!( + fs::read_to_string(repo_dir.join("note.txt")).expect("read note"), + "updated\n" + ); +} + +#[tokio::test] +async fn files_tool_lists_and_searches_with_staged_state() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-list-search", + "files", + r#"{"operations":[{"op":"create","path":"src/new.txt","content":"beta\ncreated\n"},{"op":"list","path":".","depth":2,"limit":10},{"op":"search","path":".","pattern":"beta","limit":10}]}"#, + ), + text_stream(&model.id, "done"), + ], + ); + + let repo_dir = temp_team_dir("files-list-search"); + fs::create_dir_all(repo_dir.join("src")).expect("create src"); + fs::write(repo_dir.join("note.txt"), "alpha\nbeta\n").expect("write note"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "create a file, then list and search".to_string(), + }]) + .await + .expect("send"); + + let tool_output = match &agent.history()[2] { + Message { + role: Role::User, + content, + } => match content.first().expect("tool result") { + ContentBlock::ToolResult { content, .. } => content.clone(), + other => panic!("unexpected content block: {other:?}"), + }, + _ => panic!("expected tool result"), + }; + assert!(tool_output.contains("create src/new.txt")); + assert!(tool_output.contains("[file] note.txt")); + assert!(tool_output.contains("[dir] src")); + assert!(tool_output.contains("[file] src/new.txt")); + assert!(tool_output.contains("note.txt:2: beta")); + assert!(tool_output.contains("src/new.txt:1: beta")); +} + +#[tokio::test] +async fn files_tool_aborts_without_partial_mutation_on_validation_error() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-fail", + "files", + r#"{"operations":[{"op":"create","path":"draft.txt","content":"hello"},{"op":"replace","path":"draft.txt","old":"missing","new":"present"}]}"#, + ), + text_stream(&model.id, "handled"), + ], + ); + + let repo_dir = temp_team_dir("files-no-partial"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "this edit should fail cleanly".to_string(), + }]) + .await + .expect("send"); + + match &agent.history()[2] { + Message { + role: Role::User, + content, + } => match content.first().expect("tool result") { + ContentBlock::ToolResult { + is_error: true, + content, + .. + } => assert!(content.contains("Expected 1 replacement(s)")), + other => panic!("unexpected content block: {other:?}"), + }, + _ => panic!("expected tool result"), + } + assert!(!repo_dir.join("draft.txt").exists()); +} + +#[tokio::test] +async fn files_tool_denies_writes_outside_workspace_roots() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-denied", + "files", + r#"{"operations":[{"op":"create","path":"../outside.txt","content":"nope"}]}"#, + ), + text_stream(&model.id, "handled"), + ], + ); + + let repo_dir = temp_team_dir("files-denied"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "try to write outside the repo".to_string(), + }]) + .await + .expect("send"); + + match &agent.history()[2] { + Message { + role: Role::User, + content, + } => match content.first().expect("tool result") { + ContentBlock::ToolResult { + is_error: true, + content, + .. + } => assert!(content.contains("outside the runtime policy write roots")), + other => panic!("unexpected content block: {other:?}"), + }, + _ => panic!("expected tool result"), + } + assert!(!repo_dir.join("..").join("outside.txt").exists()); +} + +#[cfg(unix)] +#[tokio::test] +async fn files_tool_search_handles_symlink_loops() { + use std::os::unix::fs::symlink; + + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "files-search-loop", + "files", + r#"{"operations":[{"op":"search","path":".","pattern":"never-match","limit":10}]}"#, + ), + text_stream(&model.id, "done"), + ], + ); + + let repo_dir = temp_team_dir("files-search-loop"); + fs::write(repo_dir.join("note.txt"), "alpha\nbeta\n").expect("write note"); + symlink(".", repo_dir.join("loop")).expect("create loop symlink"); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + timeout( + Duration::from_secs(2), + agent.send(vec![ContentBlock::Text { + text: "search the repo".to_string(), + }]), + ) + .await + .expect("search should finish") + .expect("send"); + + match &agent.history()[2] { + Message { + role: Role::User, + content, + } => match content.first().expect("tool result") { + ContentBlock::ToolResult { + is_error: false, + content, + .. + } => assert!(content.contains("(no matches)")), + other => panic!("unexpected content block: {other:?}"), + }, + _ => panic!("expected tool result"), + } +} + +#[cfg(unix)] +#[tokio::test] +async fn files_tool_list_rejects_symlink_escape_during_recursive_traversal() { + assert_files_recursive_operation_rejects_symlink_escape( + r#"{"op":"list","path":".","depth":2,"limit":10}"#, + "list", + ) + .await; +} + +#[cfg(unix)] +#[tokio::test] +async fn files_tool_search_rejects_symlink_escape_during_recursive_traversal() { + assert_files_recursive_operation_rejects_symlink_escape( + r#"{"op":"search","path":".","pattern":"outside-secret","limit":10}"#, + "search", + ) + .await; +} + +#[cfg(unix)] +async fn assert_files_recursive_operation_rejects_symlink_escape(operation: &str, label: &str) { + use std::os::unix::fs::symlink; + + let model = model_info("model", BuiltinProvider::Anthropic); + let input = format!(r#"{{"operations":[{operation}]}}"#); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + &format!("files-{label}-symlink-escape"), + "files", + &input, + ), + text_stream(&model.id, "handled"), + ], + ); + + let repo_dir = temp_team_dir(&format!("files-{label}-symlink-root")); + let outside_dir = temp_team_dir(&format!("files-{label}-symlink-outside")); + fs::write(outside_dir.join("secret.txt"), "outside-secret\n").expect("write outside file"); + symlink(&outside_dir, repo_dir.join("escape")).expect("create escape symlink"); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("inspect the repo")]) + .await + .expect("send"); + + match &agent.history()[2] { + Message { + role: Role::User, + content, + } => match content.first().expect("tool result") { + ContentBlock::ToolResult { + is_error: true, + content, + .. + } => assert!( + content.contains("outside the runtime policy read roots"), + "unexpected denial: {content}" + ), + other => panic!("unexpected content block: {other:?}"), + }, + other => panic!("expected tool result, got {other:?}"), + } + + fs::remove_dir_all(&repo_dir).expect("remove workspace"); + fs::remove_dir_all(&outside_dir).expect("remove outside directory"); +} + +#[test] +fn builtin_file_tool_profiles_default_to_batched_and_reconfigure_eagerly() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model], Vec::new()); + + let default_runtime = Runtime::builder() + .with_provider_instance(provider.clone()) + .build() + .expect("build default runtime"); + let default_names = default_runtime + .tools() + .into_iter() + .map(|tool| tool.provider.name) + .collect::>(); + assert!(default_names.contains("files")); + assert!(!default_names.contains("read")); + + let split_runtime = Runtime::builder() + .with_file_tools(FileToolProfile::Both) + .with_file_tools(FileToolProfile::Split) + .with_provider_instance(provider.clone()) + .build() + .expect("build split runtime"); + let split_names = split_runtime + .tools() + .into_iter() + .map(|tool| tool.provider.name) + .collect::>(); + assert!(!split_names.contains("files")); + for name in ["read", "ls", "grep", "glob", "write", "edit"] { + assert!(split_names.contains(name), "missing {name}"); + } + + let both_runtime = Runtime::builder() + .with_file_tools(FileToolProfile::Split) + .with_file_tools(FileToolProfile::Both) + .with_provider_instance(provider) + .build() + .expect("build combined runtime"); + let both_names = both_runtime + .tools() + .into_iter() + .map(|tool| tool.provider.name) + .collect::>(); + assert!(both_names.contains("files")); + for name in ["read", "ls", "grep", "glob", "write", "edit"] { + assert!(both_names.contains(name), "missing {name}"); + } +} + +#[tokio::test] +async fn split_read_list_and_grep_match_equivalent_files_batches() { + let repo_dir = temp_team_dir("split-read-parity"); + fs::create_dir_all(repo_dir.join("src")).expect("create src"); + fs::write(repo_dir.join("note.txt"), "alpha\nbeta\n").expect("write note"); + fs::write(repo_dir.join("src/lib.rs"), "fn beta() {}\n").expect("write source"); + + let cases = [ + ( + "read", + json!({ "path": "note.txt", "offset": 2, "limit": 1 }), + json!({ "operations": [{ "op": "read", "path": "note.txt", "offset": 2, "limit": 1 }] }), + ), + ( + "ls", + json!({ "path": ".", "depth": 2, "limit": 20 }), + json!({ "operations": [{ "op": "list", "path": ".", "depth": 2, "limit": 20 }] }), + ), + ( + "grep", + json!({ "path": ".", "pattern": "beta", "limit": 20 }), + json!({ "operations": [{ "op": "search", "path": ".", "pattern": "beta", "limit": 20 }] }), + ), + ]; + + for (tool, split_input, batched_input) in cases { + let split = + run_builtin_file_tool(FileToolProfile::Split, tool, split_input, &repo_dir).await; + let batched = + run_builtin_file_tool(FileToolProfile::Batched, "files", batched_input, &repo_dir) + .await; + + assert_eq!( + first_tool_result(&split.0), + first_tool_result(&batched.0), + "{tool} must project the same workspace-engine result as files" + ); + } + + fs::remove_dir_all(repo_dir).expect("remove workspace"); +} + +#[tokio::test] +async fn split_write_and_edit_match_equivalent_files_mutations() { + let batched_dir = temp_team_dir("batched-mutation-parity"); + let split_dir = temp_team_dir("split-mutation-parity"); + fs::write(batched_dir.join("note.txt"), "before\n").expect("write batched note"); + fs::write(split_dir.join("note.txt"), "before\n").expect("write split note"); + + run_builtin_file_tool( + FileToolProfile::Batched, + "files", + json!({ "operations": [{ "op": "set", "path": "note.txt", "content": "alpha beta\n" }] }), + &batched_dir, + ) + .await; + run_builtin_file_tool( + FileToolProfile::Split, + "write", + json!({ "filePath": "note.txt", "content": "alpha beta\n" }), + &split_dir, + ) + .await; + assert_eq!( + fs::read(batched_dir.join("note.txt")).expect("read batched write"), + fs::read(split_dir.join("note.txt")).expect("read split write") + ); + + run_builtin_file_tool( + FileToolProfile::Batched, + "files", + json!({ "operations": [{ "op": "replace", "path": "note.txt", "old": "beta", "new": "gamma" }] }), + &batched_dir, + ) + .await; + let (split_agent, _) = run_builtin_file_tool( + FileToolProfile::Split, + "edit", + json!({ + "file_path": "note.txt", + "old_string": "beta", + "new_string": "gamma" + }), + &split_dir, + ) + .await; + assert_eq!( + fs::read(batched_dir.join("note.txt")).expect("read batched edit"), + fs::read(split_dir.join("note.txt")).expect("read split edit") + ); + let detail = split_agent + .transcript() + .items() + .iter() + .find_map(|item| item.detail("call-edit")) + .expect("edit details"); + assert!(detail["diff"].as_str().is_some_and(|diff| !diff.is_empty())); + assert!( + detail["patch"] + .as_str() + .is_some_and(|patch| !patch.is_empty()) + ); + assert_eq!(detail["first_changed_line"], json!(1)); + + fs::remove_dir_all(batched_dir).expect("remove batched workspace"); + fs::remove_dir_all(split_dir).expect("remove split workspace"); +} + +#[tokio::test] +async fn every_split_file_tool_enforces_its_workspace_policy() { + let repo_dir = temp_team_dir("split-policy-root"); + let outside_dir = temp_team_dir("split-policy-outside"); + fs::write(outside_dir.join("secret.txt"), "outside-secret\n").expect("write secret"); + let outside = outside_dir.to_string_lossy().into_owned(); + let secret = outside_dir + .join("secret.txt") + .to_string_lossy() + .into_owned(); + let denied_write = outside_dir.join("new.txt").to_string_lossy().into_owned(); + let cases = [ + ("read", json!({ "path": secret }), "read roots"), + ("ls", json!({ "path": outside.clone() }), "read roots"), + ( + "grep", + json!({ "path": outside.clone(), "pattern": "secret" }), + "read roots", + ), + ( + "glob", + json!({ "path": outside, "pattern": "**/*.txt" }), + "read roots", + ), + ( + "write", + json!({ "path": denied_write, "content": "denied" }), + "write roots", + ), + ( + "edit", + json!({ + "path": outside_dir.join("secret.txt"), + "edits": [{ "old_string": "outside-secret", "new_string": "changed" }] + }), + "write roots", + ), + ]; + + for (tool, input, expected) in cases { + let (agent, _) = + run_builtin_file_tool(FileToolProfile::Split, tool, input, &repo_dir).await; + let (content, is_error) = first_tool_result(&agent); + assert!(is_error, "{tool} should be denied: {content}"); + assert!( + content.contains(expected), + "unexpected {tool} denial: {content}" + ); + } + + fs::remove_dir_all(repo_dir).expect("remove workspace"); + fs::remove_dir_all(outside_dir).expect("remove outside directory"); +} + +#[tokio::test] +async fn split_read_only_batch_preserves_provider_call_order() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-read", "read", r#"{"path":"note.txt"}"#), + ("call-glob", "glob", r#"{"pattern":"**/*.txt"}"#), + ("call-ls", "ls", r#"{"path":"."}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let repo_dir = temp_team_dir("split-parallel-order"); + fs::write(repo_dir.join("note.txt"), "alpha\n").expect("write note"); + let runtime = Runtime::builder() + .with_file_tools(FileToolProfile::Split) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("inspect in parallel")]) + .await + .expect("send"); + + let result_ids = agent.history()[2] + .content + .iter() + .filter_map(|block| match block { + ContentBlock::ToolResult { tool_use_id, .. } => Some(tool_use_id.as_str()), + _ => None, + }) + .collect::>(); + assert_eq!(result_ids, ["call-read", "call-glob", "call-ls"]); + fs::remove_dir_all(repo_dir).expect("remove workspace"); +} + +#[tokio::test] +async fn run_options_tool_budget_blocks_second_tool_call() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![multi_tool_use_stream( + &model.id, + &[ + ( + "read-1", + "files", + r#"{"operations":[{"op":"read","path":"note.txt"}]}"#, + ), + ( + "read-2", + "files", + r#"{"operations":[{"op":"read","path":"note.txt"}]}"#, + ), + ], + )], + ); + + let repo_dir = temp_team_dir("tool-budget"); + fs::write(repo_dir.join("note.txt"), "alpha\nbeta\n").expect("write note"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + let error = agent + .run( + vec![ContentBlock::Text { + text: "read the file twice".to_string(), + }], + RunOptions { + tool_budget: Some(1), + ..RunOptions::default() + }, + ) + .await + .expect_err("tool budget should fail"); + + assert!(matches!(error, RuntimeError::ToolBudgetExceeded(1))); + assert!(agent.history().is_empty()); +} + +#[tokio::test] +async fn parallel_tools_run_concurrently_and_preserve_history_order() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-1", "probe_one", r#"{}"#), + ("call-2", "probe_two", r#"{}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let log = Arc::new(tokio::sync::Mutex::new(Vec::new())); + let active = Arc::new(AtomicUsize::new(0)); + let max_active = Arc::new(AtomicUsize::new(0)); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(ProbeTool::new( + "probe_one", + true, + Duration::from_millis(40), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(ProbeTool::new( + "probe_two", + true, + Duration::from_millis(40), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run the probes")]) + .await + .expect("send"); + + assert!(max_active.load(Ordering::SeqCst) >= 2); + let tool_result_ids = agent + .history() + .iter() + .filter(|message| message.role == Role::User) + .flat_map(|message| { + message.content.iter().filter_map(|block| match block { + ContentBlock::ToolResult { tool_use_id, .. } => Some(tool_use_id.clone()), + _ => None, + }) + }) + .collect::>(); + assert_eq!( + tool_result_ids, + vec!["call-1".to_string(), "call-2".to_string()] + ); +} + +#[tokio::test] +async fn parallel_batches_respect_exclusive_barriers() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-1", "probe_one", r#"{}"#), + ("call-2", "probe_two", r#"{}"#), + ("call-3", "exclusive_probe", r#"{}"#), + ("call-4", "probe_three", r#"{}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let log = Arc::new(tokio::sync::Mutex::new(Vec::new())); + let active = Arc::new(AtomicUsize::new(0)); + let max_active = Arc::new(AtomicUsize::new(0)); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(ProbeTool::new( + "probe_one", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(ProbeTool::new( + "probe_two", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(ProbeTool::new( + "exclusive_probe", + false, + Duration::from_millis(5), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(ProbeTool::new( + "probe_three", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run probes with barrier")]) + .await + .expect("send"); + + let log = log.lock().await.clone(); + let exclusive_start = log + .iter() + .position(|entry| entry == "exclusive_probe:start") + .expect("exclusive start"); + let probe_one_start = log + .iter() + .position(|entry| entry == "probe_one:start") + .expect("probe one start"); + let probe_two_start = log + .iter() + .position(|entry| entry == "probe_two:start") + .expect("probe two start"); + let probe_three_start = log + .iter() + .position(|entry| entry == "probe_three:start") + .expect("probe three start"); + + assert!(probe_one_start < exclusive_start); + assert!(probe_two_start < exclusive_start); + assert!(exclusive_start < probe_three_start); + assert!(max_active.load(Ordering::SeqCst) >= 2); +} + +#[tokio::test] +async fn cancellation_during_parallel_batch_rolls_back_run() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![multi_tool_use_stream( + &model.id, + &[ + ("call-1", "probe_one", r#"{}"#), + ("call-2", "probe_two", r#"{}"#), + ], + )], + ); + let log = Arc::new(tokio::sync::Mutex::new(Vec::new())); + let active = Arc::new(AtomicUsize::new(0)); + let max_active = Arc::new(AtomicUsize::new(0)); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(ProbeTool::new( + "probe_one", + true, + Duration::from_millis(150), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(ProbeTool::new( + "probe_two", + true, + Duration::from_millis(150), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let cancellation = CancellationToken::default(); + let cancellation_clone = cancellation.clone(); + tokio::spawn(async move { + sleep(Duration::from_millis(20)).await; + cancellation_clone.cancel(); + }); + + let error = agent + .run( + vec![ContentBlock::text("cancel while probes run")], + RunOptions { + cancellation: Some(cancellation), + ..Default::default() + }, + ) + .await + .expect_err("run should cancel"); + + assert!(matches!(error, RuntimeError::Cancelled)); + assert!(agent.history().is_empty()); +} + +#[tokio::test] +async fn run_options_model_budget_blocks_follow_up_round() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "read-1", + "files", + r#"{"operations":[{"op":"read","path":"note.txt"}]}"#, + ), + text_stream(&model.id, "done"), + ], + ); + + let repo_dir = temp_team_dir("model-budget"); + fs::write(repo_dir.join("note.txt"), "alpha\nbeta\n").expect("write note"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(&repo_dir), + ..Default::default() + }, + ) + .expect("spawn agent"); + + let error = agent + .run( + vec![ContentBlock::Text { + text: "read the file".to_string(), + }], + RunOptions { + model_budget: Some(1), + ..Default::default() + }, + ) + .await + .expect_err("model budget should fail"); + + assert!(matches!(error, RuntimeError::ModelBudgetExceeded(1))); + assert!(agent.history().is_empty()); +} + +#[tokio::test] +async fn run_options_cancelled_run_stops_before_provider_request() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "done")], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let cancellation = CancellationToken::default(); + cancellation.cancel(); + + let error = agent + .run( + vec![ContentBlock::Text { + text: "stop".to_string(), + }], + RunOptions { + cancellation: Some(cancellation), + ..Default::default() + }, + ) + .await + .expect_err("cancellation should fail"); + + assert!(matches!(error, RuntimeError::Cancelled)); + assert!(provider_handle.recorded_requests().await.is_empty()); + assert!(agent.history().is_empty()); +} + +#[tokio::test] +async fn run_options_stop_after_tool_round_commits_transcript_and_halts() { + // A graceful stop tripped from inside a tool round ends the run at the next + // round boundary, COMMITTING the gathered transcript (unlike `cancellation`, + // which rolls it back) and making no further model request — so a follow-up + // turn on the same agent sees the gathered context. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "tool-1".to_string(), + name: "stop_tool".to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(r#"{"value":"x"}"#.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + // A second response that must NOT be consumed: the stop halts the run + // before this round's model request is issued. + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-2".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("must not run".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]), + ], + ); + let provider_handle = provider.clone(); + let stop = CancellationToken::default(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StopTrippingTool::new("stop_tool", stop.clone())) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let result = agent + .run( + vec![ContentBlock::Text { + text: "go".to_string(), + }], + RunOptions { + stop: Some(stop), + ..Default::default() + }, + ) + .await; + + // Stopped after the tool round before any final answer: an honest stop, not a + // failure — the gathered transcript is committed (user, the assistant tool + // call, the tool result), not rolled back the way cancellation would. + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + assert_eq!( + agent.history().len(), + 3, + "the gathered round is preserved, not rolled back" + ); + assert_eq!( + provider_handle.recorded_requests().await.len(), + 1, + "the stop halted the run before the second model request" + ); +} + +#[tokio::test] +async fn set_reasoning_updates_the_configured_reasoning_for_future_turns() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + // Default: no reasoning configured (the provider's default effort). + assert_eq!(agent.config().provider_request_options.reasoning, None); + + // set_reasoning sets it for future turns (mirrors set_model), enabling the + // per-phase effort split (gather low, synthesis high) on a single agent. + let high = crate::provider::ReasoningOptions { + effort: Some(crate::provider::ReasoningEffort::High), + summary: None, + }; + agent + .set_reasoning(Some(high.clone())) + .expect("set reasoning"); + assert_eq!( + agent.config().provider_request_options.reasoning, + Some(high) + ); + + // None clears it back to the provider default. + agent.set_reasoning(None).expect("clear reasoning"); + assert_eq!(agent.config().provider_request_options.reasoning, None); +} + +#[tokio::test] +async fn completed_background_results_are_batched_in_completion_order() { + let first_command = background_success_command("first", 20); + let second_command = background_success_command("second", 50); + let first_input = command_input_json(&first_command); + let second_input = command_input_json(&second_command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("tool-bg-1", "background_run", first_input.as_str()), + ("tool-bg-2", "background_run", second_input.as_str()), + ], + ), + text_stream(&model.id, "continued"), + text_stream(&model.id, "next turn"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_store(temp_store("bg-results-batched-order")) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "run two background commands".to_string(), + }]) + .await + .unwrap(); + wait_for_background_task_count(&agent, 2).await; + wait_for_background_tasks(&agent, 2, BackgroundTaskStatus::Finished).await; + + agent + .send(vec![ContentBlock::Text { + text: "report completions".to_string(), + }]) + .await + .unwrap(); + + let requests = provider_handle.recorded_requests().await; + let injected = latest_background_results_text(&requests[2]).expect("background results"); + let first = injected.find("[bg:bg-1]").expect("first task line"); + let second = injected.find("[bg:bg-2]").expect("second task line"); + assert!(first < second); + assert!(injected.contains("output=\"first")); + assert!(injected.contains("output=\"second")); +} + +#[tokio::test] +async fn failed_background_results_surface_in_snapshot_events_and_notifications() { + let command = background_failure_command("boom", 7, 50); + let input = command_input_json(&command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-bg", "background_run", &input), + text_stream(&model.id, "continued"), + text_stream(&model.id, "next turn"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_store(temp_store("bg-failure-results")) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "run failing background command".to_string(), + }]) + .await + .unwrap(); + wait_for_background_tasks(&agent, 1, BackgroundTaskStatus::Failed).await; + + let background_tasks = agent.watch_snapshot().borrow().background_tasks.clone(); + assert_eq!(background_tasks.len(), 1); + assert_eq!(background_tasks[0].status, BackgroundTaskStatus::Failed); + assert!( + background_tasks[0] + .output_preview + .as_deref() + .is_some_and(|preview| preview.contains("boom")) + ); + + agent + .send(vec![ContentBlock::Text { + text: "report failure".to_string(), + }]) + .await + .unwrap(); + + let events = collect_events(&mut events); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::BackgroundTaskFinished { task } + if task.id == "bg-1" + && task.status == BackgroundTaskStatus::Failed + && task.output_preview.as_deref().is_some_and(|preview| preview.contains("boom")) + ))); + + let requests = provider_handle.recorded_requests().await; + let injected = latest_background_results_text(&requests[2]).expect("background results"); + assert!(injected.contains("status=failed")); + assert!(injected.contains("output=\"boom")); +} + +#[tokio::test] +async fn drained_background_notifications_are_requeued_after_failed_run() { + let command = background_success_command("bg-done", 50); + let input = command_input_json(&command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-bg", "background_run", &input), + text_stream(&model.id, "continued"), + erroring_stream( + vec![ProviderEvent::MessageStarted { + id: "msg-fail".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }], + ProviderError::MalformedStream("boom".to_string()), + ), + text_stream(&model.id, "retried"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_store(temp_store("bg-requeue-failed-run")) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "run background command".to_string(), + }]) + .await + .unwrap(); + wait_for_background_tasks(&agent, 1, BackgroundTaskStatus::Finished).await; + + let failed = agent + .send(vec![ContentBlock::Text { + text: "this turn fails".to_string(), + }]) + .await; + assert!(failed.is_err()); + + agent + .send(vec![ContentBlock::Text { + text: "retry".to_string(), + }]) + .await + .unwrap(); + + let requests = provider_handle.recorded_requests().await; + let failed_request_results = + latest_background_results_text(&requests[2]).expect("background results on failed run"); + let retried_request_results = + latest_background_results_text(&requests[3]).expect("background results on retried run"); + assert_eq!(failed_request_results, retried_request_results); +} + +#[tokio::test] +async fn default_runtime_exposes_task_and_new_empty_does_not() { + let model = model_info("model", BuiltinProvider::Anthropic); + + let default_provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "ok")], + ); + let default_handle = default_provider.clone(); + let default_runtime = Runtime::builder() + .with_provider_instance(default_provider) + .build() + .expect("build runtime"); + let mut default_agent = default_runtime.spawn("agent", model.clone()).unwrap(); + default_agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await + .unwrap(); + + let default_requests = default_handle.recorded_requests().await; + let default_tools = tool_names(&default_requests[0]); + assert!(default_tools.contains("shell")); + assert!(default_tools.contains("background_run")); + assert!(default_tools.contains("check_background")); + assert!(default_tools.contains("compact")); + assert!(default_tools.contains("memory_search")); + assert!(default_tools.contains("memory_pin")); + assert!(default_tools.contains("memory_forget")); + assert!(default_tools.contains("files")); + assert!(!default_tools.contains("read_file")); + assert!(default_tools.contains("task")); + assert!(default_tools.contains("task_create")); + assert!(default_tools.contains("task_claim")); + assert!(default_tools.contains("task_update")); + assert!(default_tools.contains("task_list")); + assert!(default_tools.contains("task_get")); + assert!(!default_tools.contains("load_skill")); + + let empty_provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "ok")], + ); + let empty_handle = empty_provider.clone(); + let empty_runtime = Runtime::empty_builder() + .with_provider_instance(empty_provider) + .build() + .expect("build runtime"); + let mut empty_agent = empty_runtime.spawn("agent", model).unwrap(); + empty_agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await + .unwrap(); + + let empty_requests = empty_handle.recorded_requests().await; + let empty_tools = tool_names(&empty_requests[0]); + assert!(!empty_tools.contains("background_run")); + assert!(!empty_tools.contains("check_background")); + assert!(!empty_tools.contains("files")); + assert!(!empty_tools.contains("compact")); + assert!(!empty_tools.contains("memory_search")); + assert!(!empty_tools.contains("memory_pin")); + assert!(!empty_tools.contains("memory_forget")); + assert!(!empty_tools.contains("task")); + assert!(!empty_tools.contains("task_create")); + assert!(!empty_tools.contains("task_claim")); + assert!(!empty_tools.contains("task_update")); + assert!(!empty_tools.contains("task_list")); + assert!(!empty_tools.contains("task_get")); + assert!(!empty_tools.contains("load_skill")); +} + +#[tokio::test] +async fn tool_profile_only_exposes_requested_tools() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "ok")], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("echo_tool", "echoed")) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + tool_profile: ToolProfile::only(["files", "echo_tool"]), + ..Default::default() + }, + ) + .unwrap(); + + agent.send(vec![ContentBlock::text("hello")]).await.unwrap(); + + let requests = provider_handle.recorded_requests().await; + let tools = tool_names(&requests[0]); + assert!(tools.contains("files")); + assert!(tools.contains("echo_tool")); + assert!(!tools.contains("shell")); + assert!(!tools.contains("task")); + assert!(!tools.contains("background_run")); +} + +#[tokio::test] +async fn tool_profile_hide_blocks_named_tools() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "ok")], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + tool_profile: ToolProfile::hide(["shell", "files"]), + ..Default::default() + }, + ) + .unwrap(); + + agent.send(vec![ContentBlock::text("hello")]).await.unwrap(); + + let requests = provider_handle.recorded_requests().await; + let tools = tool_names(&requests[0]); + assert!(!tools.contains("shell")); + assert!(!tools.contains("files")); + assert!(tools.contains("task")); + assert!(tools.contains("memory_search")); +} + +#[tokio::test] +async fn hidden_tool_profile_tool_choice_falls_back_to_auto() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "ok")], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + tool_choice: Some(ToolChoice::Tool { + name: "shell".to_string(), + }), + tool_profile: ToolProfile::hide(["shell"]), + ..Default::default() + }, + ) + .unwrap(); + + agent.send(vec![ContentBlock::text("hello")]).await.unwrap(); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests[0].tool_choice, Some(ToolChoice::Auto)); + assert!(!tool_names(&requests[0]).contains("shell")); +} + +#[tokio::test] +async fn hidden_tool_profile_filters_deferred_tools_before_provider_request() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "ok")], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_tool(StaticTool::deferred_success("deferred_tool", "echoed")) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + tool_profile: ToolProfile::hide(["deferred_tool"]), + provider_request_options: crate::provider::ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Hosted, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + + agent.send(vec![ContentBlock::text("hello")]).await.unwrap(); + + let requests = provider_handle.recorded_requests().await; + assert_eq!( + requests[0].provider_request_options.tool_search_mode, + ToolSearchMode::Hosted + ); + assert!(!tool_names(&requests[0]).contains("deferred_tool")); +} + +#[tokio::test] +async fn deferred_tool_choice_reaches_provider_request_unchanged() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "ok")], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_tool(StaticTool::deferred_success("deferred_tool", "echoed")) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + tool_choice: Some(ToolChoice::Tool { + name: "deferred_tool".to_string(), + }), + provider_request_options: crate::provider::ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Hosted, + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + + agent.send(vec![ContentBlock::text("hello")]).await.unwrap(); + + let requests = provider_handle.recorded_requests().await; + assert_eq!( + requests[0].tool_choice, + Some(ToolChoice::Tool { + name: "deferred_tool".to_string(), + }) + ); + let deferred_tool = requests[0] + .tools + .iter() + .find(|tool| tool.name == "deferred_tool") + .expect("deferred tool present"); + assert_eq!( + deferred_tool.loading_policy, + crate::tool::ToolLoadingPolicy::Deferred + ); + assert_eq!( + requests[0].provider_request_options.tool_search_mode, + ToolSearchMode::Hosted + ); +} + +#[tokio::test] +async fn memory_search_tool_returns_provenance_fields() { + let store = hybrid_temp_store("memory-search-tool"); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "memory-search-1", + "memory_search", + r#"{"query":"short answers","limit":5}"#, + ), + text_stream(&model.id, "searched"), + ], + ); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let agent_id = agent.id().to_string(); + store + .upsert_records(&[MemoryRecord { + record_id: "fact:search:1".to_string(), + agent_id: agent_id.clone(), + kind: MemoryRecordKind::Fact, + content: "The user likes short answers.".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("manual_pin".to_string()), + pinned: true, + score: None, + }]) + .expect("seed records"); + + agent + .send(vec![ContentBlock::Text { + text: "Search memory.".to_string(), + }]) + .await + .expect("run"); + + let result = match &agent.history()[2].content[0] { + ContentBlock::ToolResult { content, .. } => { + let raw = content.as_str(); + serde_json::from_str::(raw).unwrap_or_else(|error| { + panic!("memory_search tool should return JSON, got {raw:?}: {error}") + }) + } + other => panic!("expected tool result, got {other:?}"), + }; + let result = result.as_array().expect("memory_search results array"); + assert_eq!(result.len(), 1); + assert_eq!(result[0]["source"].as_str(), Some("manual_pin")); + assert!(result[0]["why_retrieved"].as_str().is_some()); +} + +#[tokio::test] +async fn memory_forget_tool_rejects_cross_agent_record_ids() { + let store = hybrid_temp_store("memory-forget-cross-agent"); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "memory-forget-cross-agent", + "memory_forget", + r#"{"record_id":"fact:shared:1"}"#, + ), + text_stream(&model.id, "handled"), + ], + ); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let owner = runtime + .spawn_with_config( + "owner", + model.clone(), + AgentConfig { + memory: MemoryConfig { + write_tools_enabled: true, + ..Default::default() + }, + ..Default::default() + }, + ) + .expect("spawn owner"); + let owner_id = owner.id().to_string(); + let mut other = runtime + .spawn_with_config( + "other", + model, + AgentConfig { + memory: MemoryConfig { + write_tools_enabled: true, + ..Default::default() + }, + ..Default::default() + }, + ) + .expect("spawn other"); + + store + .upsert_records(&[MemoryRecord { + record_id: "fact:shared:1".to_string(), + agent_id: owner_id.clone(), + kind: MemoryRecordKind::Fact, + content: "Owner memory only".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("manual_pin".to_string()), + pinned: true, + score: None, + }]) + .expect("seed records"); + + other + .send(vec![ContentBlock::Text { + text: "Forget that.".to_string(), + }]) + .await + .expect("run"); + + let result = match &other.history()[2].content[0] { + ContentBlock::ToolResult { content, .. } => content.as_str(), + other => panic!("expected tool result, got {other:?}"), + }; + assert!(result.contains("was not found for this agent")); + let records = store + .search_records(&owner_id, "Owner memory", 10) + .expect("search records"); + assert_eq!(records.len(), 1); +} + +#[tokio::test] +async fn registered_skills_are_exposed_and_load_skill_returns_wrapped_content() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tool-skill", "load_skill", r#"{"name":"git"}"#), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + + let skills_dir = temp_skills_dir("load-skill"); + write_skill( + &skills_dir, + "git", + "---\nname: git\ndescription: Git workflow helpers\n---\nUse feature branches.\nRun tests first.\n", + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_skills_dir(&skills_dir) + .expect("register skills") + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + system: Some("Base system prompt".to_string()), + ..Default::default() + }, + ) + .unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "hello".to_string(), + }]) + .await + .unwrap(); + + assert_eq!( + agent.history()[2], + Message::user(ContentBlock::ToolResult { + tool_use_id: "tool-skill".to_string(), + content: "\nUse feature branches.\nRun tests first.\n" + .into(), + is_error: false, + }) + ); + + let requests = provider_handle.recorded_requests().await; + let tools = tool_names(&requests[0]); + assert!(tools.contains("load_skill")); + assert_eq!( + requests[0].system.as_deref(), + Some( + "Base system prompt\n\nSkills available:\n - git: Git workflow helpers\nUse the load_skill tool only when one of these skills is relevant to the task." + ) + ); +} + +#[tokio::test] +async fn task_subagent_keeps_load_skill_while_hiding_task() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "tool-parent", + "task", + r#"{"prompt":"inspect repo"}"#, + ), + text_stream(&model.id, "child summary"), + text_stream(&model.id, "parent done"), + ], + ); + let provider_handle = provider.clone(); + + let skills_dir = temp_skills_dir("subagent-skills"); + write_skill( + &skills_dir, + "review", + "---\nname: review\ndescription: Code review checklist\n---\nCheck tests.\n", + ); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_skills_dir(&skills_dir) + .expect("register skills") + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "delegate".to_string(), + }]) + .await + .unwrap(); + + let requests = provider_handle.recorded_requests().await; + let child_tools = tool_names(&requests[1]); + assert!(child_tools.contains("load_skill")); + assert!(!child_tools.contains("task")); + assert_eq!( + requests[1].system.as_deref(), + Some( + "You are a subagent working for another agent. Solve the delegated task, use tools when helpful, and finish with a concise final answer for the parent agent.\n\nSkills available:\n - review: Code review checklist\nUse the load_skill tool only when one of these skills is relevant to the task." + ) + ); +} + +#[tokio::test] +async fn task_tool_runs_child_with_isolated_history_and_filtered_tools() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "tool-parent", + "task", + r#"{"prompt":"inspect repo"}"#, + ), + text_stream(&model.id, "child summary"), + text_stream(&model.id, "parent done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "delegate".to_string(), + }]) + .await + .unwrap(); + + assert_eq!(agent.history().len(), 6); + assert!(agent.history().iter().any(|message| { + *message + == Message::user(ContentBlock::ToolResult { + tool_use_id: "tool-parent".to_string(), + content: "child summary".into(), + is_error: false, + }) + })); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("parent done"))) + ); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 3); + assert_eq!(requests[1].messages.len(), 1); + assert_eq!(requests[1].messages[0].role, Role::User); + assert_eq!( + requests[1].messages[0].content, + vec![ContentBlock::Text { + text: "inspect repo".to_string(), + }] + ); + + let child_tools = tool_names(&requests[1]); + assert!(child_tools.contains("shell")); + assert!(child_tools.contains("files")); + assert!(!child_tools.contains("idle")); + assert!(!child_tools.contains("task")); + + let subagents = agent.watch_snapshot().borrow().subagents.clone(); + assert_eq!(subagents.len(), 1); + assert_eq!(subagents[0].name, "agent::task"); + assert_eq!(subagents[0].model, model.id); + assert_eq!(subagents[0].status, SpawnedAgentStatus::Finished); + + let events = collect_events(&mut events); + assert!( + events + .iter() + .any(|event| matches!(event, AgentEvent::SubagentSpawned { .. })) + ); + assert!( + events + .iter() + .any(|event| matches!(event, AgentEvent::SubagentFinished { .. })) + ); +} + +#[tokio::test] +async fn task_subagent_inherits_tool_profile_and_internal_task_hide() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "tool-parent", + "task", + r#"{"prompt":"inspect repo"}"#, + ), + text_stream(&model.id, "child summary"), + text_stream(&model.id, "parent done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + tool_profile: ToolProfile::only(["task", "shell", "files"]), + ..Default::default() + }, + ) + .unwrap(); + + agent + .send(vec![ContentBlock::text("delegate")]) + .await + .unwrap(); + + let requests = provider_handle.recorded_requests().await; + let parent_tools = tool_names(&requests[0]); + assert!(parent_tools.contains("task")); + assert!(parent_tools.contains("shell")); + assert!(parent_tools.contains("files")); + assert!(!parent_tools.contains("background_run")); + + let child_tools = tool_names(&requests[1]); + assert!(child_tools.contains("shell")); + assert!(child_tools.contains("files")); + assert!(!child_tools.contains("task")); + assert!(!child_tools.contains("background_run")); +} + +#[tokio::test] +async fn task_subagent_does_not_force_hidden_task_tool_choice() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "tool-parent", + "task", + r#"{"prompt":"inspect repo"}"#, + ), + text_stream(&model.id, "child summary"), + text_stream(&model.id, "parent done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + tool_choice: Some(ToolChoice::Tool { + name: "task".to_string(), + }), + ..Default::default() + }, + ) + .unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "delegate".to_string(), + }]) + .await + .unwrap(); + + let requests = provider_handle.recorded_requests().await; + assert_eq!( + requests[0].tool_choice, + Some(ToolChoice::Tool { + name: "task".to_string(), + }) + ); + assert_eq!(requests[1].tool_choice, Some(ToolChoice::Auto)); + assert!(!tool_names(&requests[1]).contains("task")); +} + +#[tokio::test] +async fn task_tool_wraps_child_failure_and_parent_continues() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "tool-parent", + "task", + r#"{"prompt":"inspect repo"}"#, + ), + erroring_stream( + vec![ProviderEvent::MessageStarted { + id: "child-msg".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }], + ProviderError::MalformedStream("boom".to_string()), + ), + text_stream(&model.id, "handled"), + ], + ); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "delegate".to_string(), + }]) + .await + .unwrap(); + + assert!(agent.history().iter().any(|message| { + *message + == Message::user(ContentBlock::ToolResult { + tool_use_id: "tool-parent".to_string(), + content: "Subagent failed: failed to stream provider response: malformed provider stream: boom" + .into(), + is_error: true, + }) + })); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("handled"))) + ); + + let subagents = agent.watch_snapshot().borrow().subagents.clone(); + assert_eq!(subagents.len(), 1); + assert!(matches!( + &subagents[0].status, + SpawnedAgentStatus::Failed(message) + if message == "failed to stream provider response: malformed provider stream: boom" + )); +} + +#[tokio::test] +async fn child_rejects_nested_task_requests_without_recursing() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "parent-task", "task", r#"{"prompt":"delegate"}"#), + tool_use_stream(&model.id, "child-task", "task", r#"{"prompt":"recurse"}"#), + text_stream(&model.id, "child recovered"), + text_stream(&model.id, "parent done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "delegate".to_string(), + }]) + .await + .unwrap(); + + assert!(agent.history().iter().any(|message| { + *message + == Message::user(ContentBlock::ToolResult { + tool_use_id: "parent-task".to_string(), + content: "child recovered".into(), + is_error: false, + }) + })); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 4); + assert!(!tool_names(&requests[1]).contains("task")); + assert_eq!(requests[2].messages.len(), 3); + assert_eq!( + requests[2].messages[2], + Message::user(ContentBlock::ToolResult { + tool_use_id: "child-task".to_string(), + content: "Tool 'task' is not available for this agent".into(), + is_error: true, + }) + ); +} + +#[tokio::test] +async fn task_tool_returns_error_when_child_hits_round_limit() { + let model = model_info("model", BuiltinProvider::Anthropic); + let mut scripts = vec![tool_use_stream( + &model.id, + "parent-task", + "task", + r#"{"prompt":"delegate"}"#, + )]; + for index in 0..30 { + scripts.push(tool_use_stream( + &model.id, + &format!("child-tool-{index}"), + "echo_tool", + r#"{"value":"ping"}"#, + )); + } + scripts.push(text_stream(&model.id, "parent handled")); + + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], scripts); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("echo_tool", "pong")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + + agent + .send(vec![ContentBlock::Text { + text: "delegate".to_string(), + }]) + .await + .unwrap(); + + assert!(agent.history().iter().any(|message| { + *message + == Message::user(ContentBlock::ToolResult { + tool_use_id: "parent-task".to_string(), + content: "Subagent failed: max rounds exceeded at 30".into(), + is_error: true, + }) + })); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("parent handled"))) + ); + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 32); +} + +#[tokio::test] +async fn team_spawn_tool_registers_persistent_teammate() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "team-spawn", + "team_spawn", + r#"{"name":"alice","role":"researcher"}"#, + ), + text_stream(&model.id, "team ready"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("spawn-tool")), + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "build a team".to_string(), + }]) + .await + .unwrap(); + + assert!(matches!( + &agent.history()[2].content[0], + ContentBlock::ToolResult { content, is_error: false, .. } + if content.contains("Spawned persistent teammate 'alice'") + )); + + let teammates = agent.watch_snapshot().borrow().teammates.clone(); + assert_eq!(teammates.len(), 1); + assert_eq!(teammates[0].name, "alice"); + assert_eq!(teammates[0].role, "researcher"); + assert_eq!(teammates[0].status, TeamMemberStatus::Idle); + + let requests = provider_handle.recorded_requests().await; + assert!(tool_names(&requests[0]).contains("team_spawn")); + assert!(tool_names(&requests[0]).contains("team_send")); + assert!(tool_names(&requests[0]).contains("team_broadcast")); + assert!(tool_names(&requests[0]).contains("team_read_inbox")); + assert!(tool_names(&requests[0]).contains("team_request")); + assert!(tool_names(&requests[0]).contains("team_respond")); + assert!(tool_names(&requests[0]).contains("team_list_requests")); + + let events = collect_events(&mut events); + assert!( + events + .iter() + .any(|event| matches!(event, AgentEvent::TeammateSpawned { teammate } if teammate.name == "alice")) + ); +} + +#[tokio::test] +async fn persistent_teammate_processes_mail_and_reports_back_to_lead() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "child-send", + "team_send", + r#"{"to":"lead","content":"investigation complete"}"#, + ), + text_stream(&model.id, "done"), + text_stream(&model.id, "thanks"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("mailbox")), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "researcher", None) + .await + .expect("spawn teammate"); + lead.send_team_message("alice", "Check the task graph") + .expect("send message"); + + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + lead.send(vec![ContentBlock::Text { + text: "status?".to_string(), + }]) + .await + .unwrap(); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 3); + let child_tools = tool_names(&requests[0]); + assert!(child_tools.contains("team_send")); + assert!(child_tools.contains("team_read_inbox")); + assert!(child_tools.contains("team_request")); + assert!(child_tools.contains("team_respond")); + assert!(child_tools.contains("team_list_requests")); + assert!(child_tools.contains("idle")); + assert!(!child_tools.contains("team_spawn")); + assert!(!child_tools.contains("team_broadcast")); + assert!(!child_tools.contains("task_create")); + assert!(child_tools.contains("task_claim")); + assert!(child_tools.contains("task_update")); + assert!(child_tools.contains("task_list")); + assert!(child_tools.contains("task_get")); + + let inbox = latest_team_inbox_text(&requests[2]).expect("team inbox"); + assert!(inbox.contains("alice")); + assert!(inbox.contains("investigation complete")); + + let teammates = lead.watch_snapshot().borrow().teammates.clone(); + assert_eq!(teammates.len(), 1); + assert_eq!(teammates[0].status, TeamMemberStatus::Idle); +} + +#[tokio::test] +async fn teammate_inherits_tool_profile_and_internal_team_hides() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "child-send", + "team_send", + r#"{"to":"lead","content":"done"}"#, + ), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("tool-profile-team")), + tool_profile: ToolProfile::only([ + "team_spawn", + "team_send", + "team_read_inbox", + "team_request", + "team_respond", + "team_list_requests", + "idle", + "task_create", + "task_claim", + "task_update", + "task_list", + "task_get", + ]), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "researcher", None) + .await + .expect("spawn teammate"); + lead.send_team_message("alice", "Check the task graph") + .expect("send message"); + + wait_for_recorded_requests(&provider_handle, 1).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + let requests = provider_handle.recorded_requests().await; + let child_tools = tool_names(&requests[0]); + assert!(child_tools.contains("team_send")); + assert!(child_tools.contains("team_read_inbox")); + assert!(child_tools.contains("team_request")); + assert!(child_tools.contains("team_respond")); + assert!(child_tools.contains("team_list_requests")); + assert!(child_tools.contains("idle")); + assert!(child_tools.contains("task_claim")); + assert!(child_tools.contains("task_update")); + assert!(child_tools.contains("task_list")); + assert!(child_tools.contains("task_get")); + assert!(!child_tools.contains("team_spawn")); + assert!(!child_tools.contains("task_create")); + assert!(!child_tools.contains("team_broadcast")); +} + +#[tokio::test] +async fn broadcast_tool_sends_to_every_other_known_agent() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "broadcast-tool", + "team_broadcast", + r#"{"content":"team sync at noon"}"#, + ), + text_stream(&model.id, "broadcasted"), + ], + ); + + let team_dir = temp_team_dir("broadcast-tool"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let alice = runtime + .spawn_with_config( + "alice", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let bob = runtime + .spawn_with_config( + "bob", + model, + AgentConfig { + team: team_config(team_dir), + ..Default::default() + }, + ) + .unwrap(); + + lead.send(vec![ContentBlock::Text { + text: "tell everyone about the sync".to_string(), + }]) + .await + .unwrap(); + + let alice_inbox = alice.read_team_inbox().expect("alice inbox"); + let bob_inbox = bob.read_team_inbox().expect("bob inbox"); + let lead_inbox = lead.read_team_inbox().expect("lead inbox"); + + assert_eq!(alice_inbox.len(), 1); + assert_eq!(alice_inbox[0].sender, "lead"); + assert_eq!(alice_inbox[0].content, "team sync at noon"); + assert_eq!(alice_inbox[0].kind, TeamMessageKind::Broadcast); + + assert_eq!(bob_inbox.len(), 1); + assert_eq!(bob_inbox[0].sender, "lead"); + assert_eq!(bob_inbox[0].content, "team sync at noon"); + assert_eq!(bob_inbox[0].kind, TeamMessageKind::Broadcast); + + assert!(lead_inbox.is_empty()); + + let tool_result = lead + .history() + .iter() + .flat_map(|message| message.content.iter()) + .find_map(|block| match block { + ContentBlock::ToolResult { content, .. } => Some(content.clone()), + _ => None, + }) + .expect("broadcast tool result"); + assert!(tool_result.contains("2 recipient(s)")); + assert!(tool_result.contains("alice")); + assert!(tool_result.contains("bob")); +} + +#[tokio::test] +async fn teammate_message_updates_lead_unread_count_before_next_turn() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "child-send", + "team_send", + r#"{"to":"lead","content":"plan is ready"}"#, + ), + text_stream(&model.id, "done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("lead-unread-message")), + ..Default::default() + }, + ) + .unwrap(); + let mut events = lead.subscribe_events(); + + lead.spawn_teammate( + "alice", + "researcher", + Some("Send me an update.".to_string()), + ) + .await + .expect("spawn teammate"); + + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_pending_team_messages(&lead, 1).await; + + assert_eq!(lead.watch_snapshot().borrow().pending_team_messages, 1); + assert_eq!(provider_handle.recorded_requests().await.len(), 2); + + let events = collect_events(&mut events); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::TeamInboxUpdated { unread_count } if *unread_count == 1 + ))); +} + +#[tokio::test] +async fn protocol_messages_update_lead_unread_count_and_clear_on_drain() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "bob-plan", + "team_request", + r#"{"to":"lead","protocol":"plan_approval","content":"risky refactor plan"}"#, + ), + text_stream(&model.id, "waiting"), + text_stream(&model.id, "lead handled it"), + text_stream(&model.id, "done waiting"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("lead-unread-protocol")), + ..Default::default() + }, + ) + .unwrap(); + let mut events = lead.subscribe_events(); + + lead.spawn_teammate( + "bob", + "refactorer", + Some("Send me a plan request.".to_string()), + ) + .await + .expect("spawn teammate"); + + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_pending_team_messages(&lead, 1).await; + + let request_id = lead.watch_snapshot().borrow().protocol_requests[0] + .request_id + .clone(); + lead.respond_team_protocol(&request_id, false, Some("too risky".to_string())) + .expect("reject plan"); + + lead.send(vec![ContentBlock::Text { + text: "review inbox".to_string(), + }]) + .await + .unwrap(); + + assert_eq!(lead.watch_snapshot().borrow().pending_team_messages, 0); + + let events = collect_events(&mut events); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::TeamInboxUpdated { unread_count } if *unread_count == 1 + ))); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::TeamInboxUpdated { unread_count } if *unread_count == 0 + ))); +} + +#[tokio::test] +async fn team_request_tool_persists_pending_request_and_updates_snapshot() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "team-request", + "team_request", + r#"{"to":"lead","protocol":"shutdown","content":"Please shut down gracefully."}"#, + ), + text_stream(&model.id, "request queued"), + ], + ); + + let team_dir = temp_team_dir("protocol-request-tool"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::Text { + text: "queue a shutdown request".to_string(), + }]) + .await + .unwrap(); + + let requests = agent.watch_snapshot().borrow().protocol_requests.clone(); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].protocol, "shutdown"); + assert_eq!(requests[0].status, TeamProtocolStatus::Pending); + assert_eq!(requests[0].to, "lead"); + + let events = collect_events(&mut events); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::TeamProtocolRequested { request } + if request.protocol == "shutdown" && request.status == TeamProtocolStatus::Pending + ))); +} + +#[tokio::test] +async fn team_respond_tool_resolves_request_and_sends_correlated_response() { + let model = model_info("model", BuiltinProvider::Anthropic); + let team_dir = temp_team_dir("protocol-respond-tool"); + let store = temp_store("protocol-respond-tool"); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![], + )) + .build() + .expect("build runtime"); + let _lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let requester = runtime + .spawn_with_config( + "reviewer", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + let request = requester + .request_team_protocol("lead", "plan_approval", "risky refactor plan") + .expect("create request"); + let request_id = request.request_id.clone(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "team-respond", + "team_respond", + &format!(r#"{{"request_id":"{request_id}","approve":true,"reason":"looks good"}}"#), + ), + text_stream(&model.id, "approved"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_store(store) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let requester = runtime + .spawn_with_config( + "reviewer", + model, + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let mut events = lead.subscribe_events(); + + lead.send(vec![ContentBlock::Text { + text: "review the plan".to_string(), + }]) + .await + .unwrap(); + + wait_for_recorded_requests(&provider_handle, 2).await; + + let protocol_requests = lead.watch_snapshot().borrow().protocol_requests.clone(); + assert_eq!(protocol_requests.len(), 1); + assert_eq!(protocol_requests[0].status, TeamProtocolStatus::Approved); + assert_eq!( + protocol_requests[0].resolution_reason.as_deref(), + Some("looks good") + ); + + let inbox = requester.read_team_inbox().expect("reviewer inbox"); + assert_eq!(inbox.len(), 1); + assert_eq!(inbox[0].kind, TeamMessageKind::Response); + assert_eq!(inbox[0].request_id.as_deref(), Some(request_id.as_str())); + assert_eq!(inbox[0].approve, Some(true)); + + let events = collect_events(&mut events); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::TeamProtocolResolved { request } + if request.request_id == request_id.as_str() + ))); +} + +#[tokio::test] +async fn team_list_requests_tool_filters_visible_requests() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "team-list", + "team_list_requests", + r#"{"status":"pending","protocol":"plan_approval","direction":"inbound"}"#, + ), + text_stream(&model.id, "listed"), + ], + ); + + let team_dir = temp_team_dir("protocol-list-tool"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let reviewer = runtime + .spawn_with_config( + "reviewer", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let architect = runtime + .spawn_with_config( + "architect", + model, + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + let pending = reviewer + .request_team_protocol("lead", "plan_approval", "plan A") + .expect("pending request"); + let resolved = architect + .request_team_protocol("lead", "shutdown", "stop") + .expect("resolved request"); + lead.respond_team_protocol(&resolved.request_id, false, Some("not now".to_string())) + .expect("resolve request"); + + lead.send(vec![ContentBlock::Text { + text: "list pending reviews".to_string(), + }]) + .await + .unwrap(); + + let tool_result = lead + .history() + .iter() + .flat_map(|message| message.content.iter()) + .find_map(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + content, + .. + } if tool_use_id == "team-list" => Some(content.to_display_string()), + _ => None, + }) + .expect("team_list_requests tool result"); + let listed: serde_json::Value = serde_json::from_str(&tool_result).unwrap_or_else(|error| { + panic!("team_list_requests should return JSON, got {tool_result:?}: {error}") + }); + let listed = listed.as_array().expect("array"); + assert_eq!(listed.len(), 1); + assert_eq!( + listed[0]["request_id"].as_str(), + Some(pending.request_id.as_str()) + ); +} + +#[tokio::test] +async fn plan_approval_request_response_keeps_teammate_alive() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "bob-plan", + "team_request", + r#"{"to":"lead","protocol":"plan_approval","content":"risky refactor plan"}"#, + ), + text_stream(&model.id, "waiting for review"), + text_stream(&model.id, "plan rejected, continuing safely"), + text_stream(&model.id, "still available"), + ], + ); + let provider_handle = provider.clone(); + + let team_dir = temp_team_dir("plan-approval"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate( + "bob", + "refactorer", + Some("Propose your plan first.".to_string()), + ) + .await + .expect("spawn teammate"); + + wait_for_recorded_requests(&provider_handle, 2).await; + let requests = lead.watch_snapshot().borrow().protocol_requests.clone(); + assert_eq!(requests.len(), 1); + let request_id = requests[0].request_id.clone(); + assert_eq!(requests[0].status, TeamProtocolStatus::Pending); + + lead.respond_team_protocol(&request_id, false, Some("too risky".to_string())) + .expect("respond to plan"); + + wait_for_recorded_requests(&provider_handle, 3).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + let requests = lead.watch_snapshot().borrow().protocol_requests.clone(); + assert_eq!(requests[0].status, TeamProtocolStatus::Rejected); + assert_eq!(requests[0].resolution_reason.as_deref(), Some("too risky")); +} + +#[tokio::test] +async fn shutdown_approval_shuts_down_teammate_after_current_wake_cycle() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (stream, tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![stream, text_stream(&model.id, "shutting down now")], + ); + let provider_handle = provider.clone(); + + let team_dir = temp_team_dir("shutdown-protocol"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "coder", None) + .await + .expect("spawn teammate"); + + let request = lead + .request_team_protocol("alice", "shutdown", "Please stop after this turn.") + .expect("shutdown request"); + + tx.send(Ok(ProviderEvent::MessageStarted { + id: "msg-shutdown".to_string(), + model: "model".to_string(), + role: Role::Assistant, + })) + .expect("message start"); + tx.send(Ok(ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "shutdown-response".to_string(), + name: "team_respond".to_string(), + }, + })) + .expect("tool start"); + tx.send(Ok(ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(format!( + r#"{{"request_id":"{}","approve":true,"reason":"wrapping up"}}"#, + request.request_id + )), + })) + .expect("tool delta"); + tx.send(Ok(ProviderEvent::ContentBlockStopped { index: 0 })) + .expect("tool stop"); + tx.send(Ok(ProviderEvent::MessageStopped)) + .expect("message stop"); + drop(tx); + + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Shutdown).await; + + let requests = lead.watch_snapshot().borrow().protocol_requests.clone(); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].status, TeamProtocolStatus::Approved); + assert_eq!( + requests[0].resolution_reason.as_deref(), + Some("wrapping up") + ); + + let teammates = lead.watch_snapshot().borrow().teammates.clone(); + assert_eq!(teammates[0].status, TeamMemberStatus::Shutdown); +} + +#[tokio::test] +async fn failed_run_requeues_protocol_messages_and_preserves_request_state() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![erroring_stream( + vec![ + ProviderEvent::MessageStarted { + id: "msg-error".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("starting".to_string()), + }, + ], + ProviderError::InvalidResponse("boom".to_string()), + )], + ); + + let team_dir = temp_team_dir("protocol-requeue"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let reviewer = runtime + .spawn_with_config( + "reviewer", + model, + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + let request = reviewer + .request_team_protocol("lead", "plan_approval", "please review") + .expect("create request"); + + let error = lead + .send(vec![ContentBlock::Text { + text: "handle inbox".to_string(), + }]) + .await + .expect_err("run should fail"); + assert!(matches!( + error, + crate::error::RuntimeError::FailedToStreamResponse(_) + )); + assert_eq!(lead.watch_snapshot().borrow().pending_team_messages, 1); + + let inbox = lead.read_team_inbox().expect("requeued inbox"); + assert_eq!(inbox.len(), 1); + assert_eq!(inbox[0].kind, TeamMessageKind::Request); + assert_eq!( + inbox[0].request_id.as_deref(), + Some(request.request_id.as_str()) + ); + + let requests = lead.watch_snapshot().borrow().protocol_requests.clone(); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].status, TeamProtocolStatus::Pending); +} + +#[tokio::test] +async fn failed_teammate_can_recover_on_next_wake() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + erroring_stream( + vec![ProviderEvent::MessageStarted { + id: "msg-error".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }], + ProviderError::InvalidResponse("boom".to_string()), + ), + text_stream(&model.id, "recovered"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("teammate-recover")), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "researcher", Some("first try".to_string())) + .await + .expect("spawn teammate"); + wait_for_teammate_status( + &lead, + TeamMemberStatus::Failed( + "failed to stream provider response: invalid provider response: boom".to_string(), + ), + ) + .await; + + lead.send_team_message("alice", "try again") + .expect("send retry"); + wait_for_recorded_requests(&provider_handle, 2).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; +} + +#[tokio::test] +async fn persisted_protocol_requests_load_on_restart() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + + let team_dir = temp_team_dir("protocol-restart"); + let store = temp_store("protocol-restart"); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + let reviewer = runtime + .spawn_with_config( + "reviewer", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + reviewer + .request_team_protocol("lead", "plan_approval", "plan one") + .expect("create request"); + + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + let runtime = Runtime::builder() + .with_store(store) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let restarted = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir), + ..Default::default() + }, + ) + .unwrap(); + + assert_eq!(lead.watch_snapshot().borrow().protocol_requests.len(), 1); + assert_eq!( + restarted.watch_snapshot().borrow().protocol_requests.len(), + 1 + ); + assert_eq!( + restarted.watch_snapshot().borrow().protocol_requests[0].status, + TeamProtocolStatus::Pending + ); + assert_eq!(restarted.watch_snapshot().borrow().pending_team_messages, 1); +} + +#[tokio::test] +async fn persisted_teammates_reload_as_shutdown_without_live_actor() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + + let team_dir = temp_team_dir("teammate-restart-status"); + let store = temp_store("teammate-restart-status"); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "researcher", None) + .await + .expect("spawn teammate"); + + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + let runtime = Runtime::builder() + .with_store(store) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let restarted = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir), + ..Default::default() + }, + ) + .unwrap(); + + let teammates = restarted.watch_snapshot().borrow().teammates.clone(); + assert_eq!(teammates.len(), 1); + assert_eq!(teammates[0].name, "alice"); + assert_eq!(teammates[0].status, TeamMemberStatus::Idle); +} + +#[tokio::test] +async fn team_spawn_revives_shutdown_teammate_name_after_restart() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + + let team_dir = temp_team_dir("teammate-revive"); + let store = temp_store("teammate-revive"); + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("bob", "researcher", None) + .await + .expect("spawn teammate"); + + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + let runtime = Runtime::builder() + .with_store(store) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut restarted = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir.clone()), + ..Default::default() + }, + ) + .unwrap(); + + restarted + .spawn_teammate("bob", "refactor specialist", None) + .await + .expect("revive teammate"); + + let teammates = restarted.watch_snapshot().borrow().teammates.clone(); + assert_eq!(teammates.len(), 1); + assert_eq!(teammates[0].name, "bob"); + assert_eq!(teammates[0].role, "refactor specialist"); + assert_eq!(teammates[0].status, TeamMemberStatus::Idle); +} + +#[tokio::test] +async fn autonomous_teammate_auto_claims_ready_task_after_spawn() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "claimed work")], + ); + let provider_handle = provider.clone(); + + let team_dir = temp_team_dir("auto-claim-team"); + let tasks_dir = temp_team_dir("auto-claim-tasks"); + let store = temp_store("auto-claim"); + create_task(&store, &tasks_dir, "Implement feature", "", vec![]); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: autonomous_team_config( + team_dir, + Duration::from_millis(10), + Duration::from_millis(120), + ), + task: TaskConfig { + tasks_dir: tasks_dir.clone(), + reminder_threshold: 3, + }, + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "coder", None) + .await + .expect("spawn teammate"); + + wait_for_recorded_requests(&provider_handle, 1).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + wait_for_snapshot_task_owner(&lead, 1, "alice").await; + + let task = load_task(&store, &tasks_dir, 1); + assert_eq!(task["owner"].as_str(), Some("alice")); + assert_eq!(lead.watch_snapshot().borrow().tasks[0].owner, "alice"); + + let requests = provider_handle.recorded_requests().await; + let auto_claim = latest_auto_claim_text(&requests[0]).expect("auto-claim text"); + assert!(auto_claim.contains("Task #1")); + assert!(auto_claim.contains("Implement feature")); + assert!(request_contains_text( + &requests[0], + "Update your task status." + )); +} + +#[tokio::test] +async fn autonomous_teammates_do_not_double_claim_same_task() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "picked it up")], + ); + let provider_handle = provider.clone(); + + let team_dir = temp_team_dir("claim-race-team"); + let tasks_dir = temp_team_dir("claim-race-tasks"); + let store = temp_store("claim-race"); + create_task(&store, &tasks_dir, "One task", "", vec![]); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: autonomous_team_config( + team_dir, + Duration::from_millis(10), + Duration::from_millis(70), + ), + task: TaskConfig { + tasks_dir: tasks_dir.clone(), + reminder_threshold: 3, + }, + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "coder", None) + .await + .expect("spawn alice"); + lead.spawn_teammate("bob", "coder", None) + .await + .expect("spawn bob"); + + wait_for_recorded_requests(&provider_handle, 1).await; + sleep(Duration::from_millis(120)).await; + + let task = load_task(&store, &tasks_dir, 1); + let owner = task["owner"].as_str().expect("owner"); + assert!(owner == "alice" || owner == "bob"); + assert_eq!(provider_handle.recorded_requests().await.len(), 1); +} + +#[tokio::test] +async fn autonomous_teammate_does_not_claim_more_work_while_owning_unfinished_task() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "looking into task 1")], + ); + let provider_handle = provider.clone(); + + let team_dir = temp_team_dir("claim-owned-work-team"); + let tasks_dir = temp_team_dir("claim-owned-work-tasks"); + let store = temp_store("claim-owned-work"); + create_task(&store, &tasks_dir, "Task one", "", vec![]); + create_task(&store, &tasks_dir, "Task two", "", vec![]); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: autonomous_team_config( + team_dir, + Duration::from_millis(10), + Duration::from_millis(80), + ), + task: TaskConfig { + tasks_dir: tasks_dir.clone(), + reminder_threshold: 3, + }, + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "coder", None) + .await + .expect("spawn teammate"); + + wait_for_recorded_requests(&provider_handle, 1).await; + wait_for_snapshot_task_owner(&lead, 1, "alice").await; + sleep(Duration::from_millis(120)).await; + + assert_eq!( + load_task(&store, &tasks_dir, 1)["owner"].as_str(), + Some("alice") + ); + assert_eq!(load_task(&store, &tasks_dir, 2)["owner"].as_str(), Some("")); + assert_eq!(provider_handle.recorded_requests().await.len(), 1); +} + +#[tokio::test] +async fn autonomous_teammate_claims_task_after_dependency_unblocks() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "starting unblocked task")], + ); + + let team_dir = temp_team_dir("unblock-team"); + let tasks_dir = temp_team_dir("unblock-tasks"); + let store = temp_store("unblock"); + create_task(&store, &tasks_dir, "Blocked elsewhere", "lead", vec![]); + create_task(&store, &tasks_dir, "Ready later", "", vec![1]); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: autonomous_team_config( + team_dir, + Duration::from_millis(10), + Duration::from_millis(300), + ), + task: TaskConfig { + tasks_dir: tasks_dir.clone(), + reminder_threshold: 3, + }, + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "coder", None) + .await + .expect("spawn teammate"); + + sleep(Duration::from_millis(30)).await; + assert_eq!(load_task(&store, &tasks_dir, 2)["owner"].as_str(), Some("")); + + task::execute_with_store( + &store, + &TaskIntrinsicTool::Update, + json!({"taskId": 1, "status": "completed"}), + tasks_dir.as_path(), + TaskAccess::Lead, + ) + .expect("complete blocker"); + + wait_for_task_owner(&store, &tasks_dir, 2, "alice").await; +} + +#[tokio::test] +async fn teammate_task_updates_are_limited_to_owned_tasks() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "own-task", + "task_update", + r#"{"taskId":1,"status":"in_progress"}"#, + ), + text_stream(&model.id, "updated own task"), + tool_use_stream( + &model.id, + "other-task", + "task_update", + r#"{"taskId":2,"status":"completed"}"#, + ), + text_stream(&model.id, "could not touch task 2"), + ], + ); + let team_dir = temp_team_dir("owned-update-team"); + let tasks_dir = temp_team_dir("owned-update-tasks"); + let store = temp_store("owned-update"); + create_task(&store, &tasks_dir, "Alice task", "alice", vec![]); + create_task(&store, &tasks_dir, "Shared task", "", vec![]); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir), + task: TaskConfig { + tasks_dir: tasks_dir.clone(), + reminder_threshold: 3, + }, + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "coder", Some("Start task 1.".to_string())) + .await + .expect("spawn teammate"); + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + lead.send_team_message("alice", "Now try to complete task 2.") + .expect("send message"); + sleep(Duration::from_millis(100)).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + assert_eq!( + load_task(&store, &tasks_dir, 1)["status"].as_str(), + Some("in_progress") + ); + assert_eq!( + load_task(&store, &tasks_dir, 2)["status"].as_str(), + Some("pending") + ); +} + +#[tokio::test] +async fn teammate_task_subagent_inherits_owner_restrictions() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "delegate-task", + "task", + r#"{"prompt":"finish the other task"}"#, + ), + tool_use_stream( + &model.id, + "illegal-update", + "task_update", + r#"{"taskId":2,"status":"completed"}"#, + ), + text_stream(&model.id, "child could not change task 2"), + text_stream(&model.id, "delegation handled"), + ], + ); + let team_dir = temp_team_dir("teammate-task-subagent"); + let tasks_dir = temp_team_dir("teammate-task-subagent-tasks"); + let store = temp_store("teammate-task-subagent"); + create_task(&store, &tasks_dir, "Alice task", "alice", vec![]); + create_task(&store, &tasks_dir, "Bob task", "bob", vec![]); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(team_dir), + task: TaskConfig { + tasks_dir: tasks_dir.clone(), + reminder_threshold: 3, + }, + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "coder", Some("delegate it".to_string())) + .await + .expect("spawn teammate"); + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + assert_eq!( + load_task(&store, &tasks_dir, 2)["status"].as_str(), + Some("pending") + ); +} + +#[tokio::test] +async fn idle_tool_returns_teammate_to_idle() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![tool_use_stream(&model.id, "idle-now", "idle", "{}")], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("idle-tool-team")), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "coder", Some("Check in then idle.".to_string())) + .await + .expect("spawn teammate"); + + wait_for_recorded_requests(&provider_handle, 1).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + assert_eq!(provider_handle.recorded_requests().await.len(), 1); +} + +#[tokio::test] +async fn autonomous_idle_timeout_shuts_down_and_same_name_can_respawn() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], vec![]); + + let team_dir = temp_team_dir("idle-timeout-team"); + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model.clone(), + AgentConfig { + team: autonomous_team_config( + team_dir.clone(), + Duration::from_millis(10), + Duration::from_millis(40), + ), + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "researcher", None) + .await + .expect("spawn teammate"); + wait_for_teammate_status(&lead, TeamMemberStatus::Shutdown).await; + + lead.spawn_teammate("alice", "researcher", None) + .await + .expect("respawn teammate"); +} + +#[tokio::test] +async fn teammate_identity_is_reinjected_after_compaction() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first done"), + text_stream(&model.id, "summary"), + text_stream(&model.id, "second done"), + text_stream(&model.id, "extra done"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = Runtime::builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut lead = runtime + .spawn_with_config( + "lead", + model, + AgentConfig { + team: team_config(temp_team_dir("identity-compact-team")), + compaction: crate::agent::CompactionConfig { + auto_compact_threshold_tokens: Some(1), + ..Default::default() + }, + ..Default::default() + }, + ) + .unwrap(); + + lead.spawn_teammate("alice", "researcher", Some("first".to_string())) + .await + .expect("spawn teammate"); + wait_for_recorded_requests(&provider_handle, 1).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + lead.send_team_message("alice", "second") + .expect("send second"); + wait_for_recorded_requests(&provider_handle, 3).await; + wait_for_teammate_status(&lead, TeamMemberStatus::Idle).await; + + let requests = provider_handle.recorded_requests().await; + assert!( + requests + .iter() + .skip(1) + .any(|request| request_contains_text(request, "")) + ); + assert!( + requests + .iter() + .skip(1) + .any(|request| request_contains_text(request, "I am alice")) + ); +} + +#[tokio::test] +async fn shell_tool_emits_progress_events_with_output_lines() { + let echo_command = "echo hello-progress"; + let input = command_input_json(echo_command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tc-shell-progress", "shell", &input), + text_stream(&model.id, "done"), + ], + ); + + let runtime = Runtime::builder() + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::text("run echo")]) + .await + .unwrap(); + + let all_events = collect_events(&mut events); + let progress_events: Vec<_> = all_events + .iter() + .filter(|event| matches!(event, AgentEvent::ToolExecutionProgress { .. })) + .collect(); + + assert!( + !progress_events.is_empty(), + "expected at least one ToolExecutionProgress event, got events: {all_events:?}" + ); + + let has_stdout = progress_events.iter().any(|event| { + if let AgentEvent::ToolExecutionProgress { progress, .. } = event { + progress.contains("hello-progress") + } else { + false + } + }); + assert!( + has_stdout, + "expected progress event containing 'hello-progress', got: {progress_events:?}" + ); + + // Verify the progress events reference the correct tool call + for event in &progress_events { + if let AgentEvent::ToolExecutionProgress { id, name, .. } = event { + assert_eq!(id, "tc-shell-progress"); + assert_eq!(name, "shell"); + } + } +} + +#[tokio::test] +async fn shell_tool_emits_stderr_progress_on_failure() { + let fail_command = if cfg!(windows) { + "echo error-output 1>&2 & exit /b 1" + } else { + "echo error-output >&2 && exit 1" + }; + let input = command_input_json(fail_command); + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "tc-stderr", "shell", &input), + text_stream(&model.id, "done"), + ], + ); + + let runtime = Runtime::builder() + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).unwrap(); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::text("run failing command")]) + .await + .unwrap(); + + let all_events = collect_events(&mut events); + let progress_events: Vec<_> = all_events + .iter() + .filter(|event| matches!(event, AgentEvent::ToolExecutionProgress { .. })) + .collect(); + + assert!( + !progress_events.is_empty(), + "expected at least one ToolExecutionProgress event for stderr, got events: {all_events:?}" + ); + + let has_stderr = progress_events.iter().any(|event| { + if let AgentEvent::ToolExecutionProgress { progress, .. } = event { + progress.contains("error-output") + } else { + false + } + }); + assert!( + has_stderr, + "expected progress event containing 'error-output', got: {progress_events:?}" + ); +} + +fn collect_events(receiver: &mut tokio::sync::broadcast::Receiver) -> Vec { + let mut events = Vec::new(); + while let Ok(event) = receiver.try_recv() { + events.push(event); + } + events +} + +const SHORT_WAIT_ATTEMPTS: usize = 200; +const BACKGROUND_WAIT_ATTEMPTS: usize = 6000; +const POLL_INTERVAL_MS: u64 = 10; + +async fn wait_for_pending_team_messages(agent: &Agent, expected_count: usize) { + for _ in 0..SHORT_WAIT_ATTEMPTS { + if agent.watch_snapshot().borrow().pending_team_messages == expected_count { + return; + } + sleep(Duration::from_millis(POLL_INTERVAL_MS)).await; + } + + panic!("timed out waiting for {expected_count} pending team messages"); +} + +async fn wait_for_background_task_count(agent: &Agent, expected_count: usize) { + for _ in 0..BACKGROUND_WAIT_ATTEMPTS { + if agent.watch_snapshot().borrow().background_tasks.len() == expected_count { + return; + } + sleep(Duration::from_millis(POLL_INTERVAL_MS)).await; + } + + panic!("timed out waiting for {expected_count} background tasks"); +} + +async fn wait_for_background_tasks( + agent: &Agent, + expected_count: usize, + status: BackgroundTaskStatus, +) { + for _ in 0..BACKGROUND_WAIT_ATTEMPTS { + let background_tasks = agent.watch_snapshot().borrow().background_tasks.clone(); + if background_tasks.len() == expected_count + && background_tasks.iter().all(|task| task.status == status) + { + return; + } + sleep(Duration::from_millis(POLL_INTERVAL_MS)).await; + } + + panic!("timed out waiting for {expected_count} background tasks to reach {status:?}"); +} + +fn latest_background_results_text<'a>(request: &'a Request<'a>) -> Option<&'a str> { + request + .messages + .iter() + .rev() + .flat_map(|message| message.content.iter()) + .find_map(|block| match block { + ContentBlock::Text { text } if text.contains("") => { + Some(text.as_str()) + } + _ => None, + }) +} + +fn latest_team_inbox_text<'a>(request: &'a Request<'a>) -> Option<&'a str> { + request + .messages + .iter() + .rev() + .flat_map(|message| message.content.iter()) + .find_map(|block| match block { + ContentBlock::Text { text } if text.contains("") => Some(text.as_str()), + _ => None, + }) +} + +fn latest_auto_claim_text<'a>(request: &'a Request<'a>) -> Option<&'a str> { + request + .messages + .iter() + .rev() + .flat_map(|message| message.content.iter()) + .find_map(|block| match block { + ContentBlock::Text { text } if text.contains("") => Some(text.as_str()), + _ => None, + }) +} + +fn request_contains_text(request: &Request<'_>, pattern: &str) -> bool { + request + .messages + .iter() + .flat_map(|message| message.content.iter()) + .any(|block| matches!(block, ContentBlock::Text { text } if text.contains(pattern))) +} + +fn request_contains_tool_result(request: &Request<'_>, pattern: &str) -> bool { + request + .messages + .iter() + .flat_map(|message| message.content.iter()) + .any(|block| { + matches!( + block, + ContentBlock::ToolResult { content, .. } if content.contains(pattern) + ) + }) +} + +fn text_stream(model: &str, text: &str) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn tool_use_stream(model: &str, id: &str, name: &str, input_json: &str) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{id}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: id.to_string(), + name: name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn multi_tool_use_stream(model: &str, calls: &[(&str, &str, &str)]) -> StreamScript { + let mut events = vec![ProviderEvent::MessageStarted { + id: "msg-multi-tool".to_string(), + model: model.to_string(), + role: Role::Assistant, + }]; + + for (index, (id, name, input_json)) in calls.iter().enumerate() { + events.push(ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::ToolUse { + id: (*id).to_string(), + name: (*name).to_string(), + }, + }); + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ToolUseInputJson((*input_json).to_string()), + }); + events.push(ProviderEvent::ContentBlockStopped { index }); + } + + events.push(ProviderEvent::MessageStopped); + ok_stream(events) +} + +async fn run_builtin_file_tool( + profile: FileToolProfile, + tool_name: &str, + input: serde_json::Value, + workspace: &Path, +) -> (Agent, ScriptedProvider) { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + &format!("call-{tool_name}"), + tool_name, + &input.to_string(), + ), + text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::builder() + .with_file_tools(profile) + .with_provider_instance(provider.clone()) + .build() + .expect("build file-tool runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: workspace_config(workspace), + ..Default::default() + }, + ) + .expect("spawn file-tool agent"); + agent + .send(vec![ContentBlock::text(format!("run {tool_name}"))]) + .await + .expect("run file tool"); + (agent, provider) +} + +fn first_tool_result(agent: &Agent) -> (String, bool) { + agent + .history() + .iter() + .flat_map(|message| message.content.iter()) + .find_map(|block| match block { + ContentBlock::ToolResult { + content, is_error, .. + } => Some((content.to_display_string(), *is_error)), + _ => None, + }) + .expect("tool result") +} + +fn tool_names(request: &Request<'_>) -> std::collections::HashSet { + request.tools.iter().map(|tool| tool.name.clone()).collect() +} + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +fn temp_skills_dir(label: &str) -> PathBuf { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = std::env::temp_dir().join(format!( + "mentra-runtime-skills-{label}-{timestamp}-{unique}" + )); + fs::create_dir_all(&path).expect("create temp dir"); + path +} + +fn temp_team_dir(label: &str) -> PathBuf { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = + std::env::temp_dir().join(format!("mentra-runtime-team-{label}-{timestamp}-{unique}")); + fs::create_dir_all(&path).expect("create team dir"); + path +} + +fn temp_store(label: &str) -> SqliteRuntimeStore { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + SqliteRuntimeStore::new(std::env::temp_dir().join(format!( + "mentra-runtime-store-{label}-{timestamp}-{unique}.sqlite" + ))) +} + +fn hybrid_temp_store(label: &str) -> HybridRuntimeStore { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let base_dir = std::env::temp_dir().join(format!( + "mentra-runtime-hybrid-store-{label}-{timestamp}-{unique}" + )); + HybridRuntimeStore::with_memory_path( + base_dir.join("runtime.sqlite"), + base_dir.join("memory.sqlite"), + ) +} + +fn team_config(team_dir: PathBuf) -> TeamConfig { + TeamConfig { + team_dir, + ..Default::default() + } +} + +fn autonomous_team_config( + team_dir: PathBuf, + poll_interval: Duration, + idle_timeout: Duration, +) -> TeamConfig { + TeamConfig { + team_dir, + autonomy: TeamAutonomyConfig { + enabled: true, + poll_interval, + idle_timeout, + }, + } +} + +fn create_task( + store: &SqliteRuntimeStore, + tasks_dir: &Path, + subject: &str, + owner: &str, + blocked_by: Vec, +) { + create_task_with_directory(store, tasks_dir, subject, owner, blocked_by, None); +} + +fn create_task_with_directory( + store: &SqliteRuntimeStore, + tasks_dir: &Path, + subject: &str, + owner: &str, + blocked_by: Vec, + working_directory: Option<&str>, +) { + task::execute_with_store( + store, + &TaskIntrinsicTool::Create, + json!({ + "subject": subject, + "owner": owner, + "workingDirectory": working_directory, + "blockedBy": blocked_by, + }), + tasks_dir, + TaskAccess::Lead, + ) + .expect("create task"); +} + +fn workspace_config(base_dir: &Path) -> WorkspaceConfig { + WorkspaceConfig { + base_dir: base_dir.to_path_buf(), + ..Default::default() + } +} + +fn load_task(store: &SqliteRuntimeStore, tasks_dir: &Path, task_id: u64) -> serde_json::Value { + serde_json::from_str( + &task::execute_with_store( + store, + &TaskIntrinsicTool::Get, + json!({ "taskId": task_id }), + tasks_dir, + TaskAccess::Lead, + ) + .expect("load task"), + ) + .expect("parse task") +} + +fn write_skill(root: &Path, name: &str, content: &str) { + let skill_dir = root.join(name); + fs::create_dir_all(&skill_dir).expect("create skill dir"); + fs::write(skill_dir.join("SKILL.md"), content).expect("write skill"); +} + +async fn wait_for_recorded_requests(provider: &ScriptedProvider, expected: usize) { + for _ in 0..BACKGROUND_WAIT_ATTEMPTS { + if provider.recorded_requests().await.len() >= expected { + return; + } + sleep(Duration::from_millis(POLL_INTERVAL_MS)).await; + } + + panic!("timed out waiting for {expected} recorded requests"); +} + +async fn wait_for_background_task_status( + store: &SqliteRuntimeStore, + agent_id: &str, + task_id: &str, + expected_status: BackgroundTaskStatus, +) { + for _ in 0..BACKGROUND_WAIT_ATTEMPTS { + let tasks = + ::load_background_tasks( + store, agent_id, + ) + .expect("load background tasks"); + if tasks + .iter() + .any(|task| task.id == task_id && task.status == expected_status) + { + return; + } + sleep(Duration::from_millis(POLL_INTERVAL_MS)).await; + } + + panic!("timed out waiting for background task {task_id} to reach {expected_status:?}"); +} + +async fn wait_for_background_task_record( + store: &SqliteRuntimeStore, + agent_id: &str, + expected_count: usize, +) { + for _ in 0..BACKGROUND_WAIT_ATTEMPTS { + let tasks = + ::load_background_tasks( + store, agent_id, + ) + .expect("load background tasks"); + if tasks.len() == expected_count { + return; + } + sleep(Duration::from_millis(POLL_INTERVAL_MS)).await; + } + + panic!("timed out waiting for {expected_count} background task records"); +} + +async fn wait_for_teammate_status(agent: &Agent, expected: TeamMemberStatus) { + for _ in 0..500 { + let teammates = agent.watch_snapshot().borrow().teammates.clone(); + if teammates.len() == 1 && teammates[0].status == expected { + return; + } + sleep(Duration::from_millis(10)).await; + } + + panic!("timed out waiting for teammate status {expected:?}"); +} + +async fn wait_for_task_owner( + store: &SqliteRuntimeStore, + tasks_dir: &Path, + task_id: u64, + owner: &str, +) { + for _ in 0..500 { + if load_task(store, tasks_dir, task_id)["owner"].as_str() == Some(owner) { + return; + } + sleep(Duration::from_millis(10)).await; + } + + panic!("timed out waiting for task {task_id} owner {owner}"); +} + +async fn wait_for_snapshot_task_owner(agent: &Agent, task_id: u64, owner: &str) { + for _ in 0..200 { + let tasks = agent.watch_snapshot().borrow().tasks.clone(); + if tasks + .iter() + .find(|task| task.id == task_id) + .map(|task| task.owner.as_str()) + == Some(owner) + { + return; + } + sleep(Duration::from_millis(10)).await; + } + + panic!("timed out waiting for snapshot task {task_id} owner {owner}"); +} + +#[tokio::test] +async fn pre_execution_hook_blocks_tool_call() { + // Define a hook that blocks "echo_tool" + struct BlockEchoHook; + #[async_trait] + impl PreExecutionHook for BlockEchoHook { + async fn pre_tool_execution( + &self, + context: &PreExecutionContext, + ) -> Result { + if context.tool_name == "echo_tool" { + Ok(HookDecision::Deny( + "echo_tool is blocked by policy".to_string(), + )) + } else { + Ok(HookDecision::Allow) + } + } + } + + // Build runtime with the hook + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + // Agent tries to call echo_tool, gets blocked error result, model responds + tool_use_stream(&model.id, "call-1", "echo_tool", r#"{"value":"hello"}"#), + text_stream(&model.id, "tool was blocked"), + ], + ); + + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("echo_tool", "should not run")) + .with_pre_hook(BlockEchoHook) + .build() + .expect("build runtime"); + + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "call the tool".to_string(), + }]) + .await + .expect("send should succeed despite blocked tool"); + + // The agent should have received the blocked error and continued with text + let history = agent.history(); + // Find the tool result in history that contains the block message + let has_blocked_result = history.iter().any(|msg| { + msg.content.iter().any(|block| { + matches!( + block, + ContentBlock::ToolResult { + content, + is_error: true, + .. + } if content.to_display_string().contains("blocked by policy") + ) + }) + }); + assert!( + has_blocked_result, + "expected blocked tool result in history" + ); +} diff --git a/vendor/mentra/src/agent/tests/runtime_volatile_store.rs b/vendor/mentra/src/agent/tests/runtime_volatile_store.rs new file mode 100644 index 0000000..a558b9f --- /dev/null +++ b/vendor/mentra/src/agent/tests/runtime_volatile_store.rs @@ -0,0 +1,427 @@ +//! Integration coverage for the volatile, no-durable-trace `RuntimeStore` +//! profile (`VolatileRuntimeStore`) against a full `Agent::run` (via +//! `Agent::send`), not just the store's own unit tests. + +use std::{ + fs, + path::PathBuf, + sync::atomic::{AtomicU64, Ordering}, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use crate::{ + BuiltinProvider, ContentBlock, Role, + agent::{AgentConfig, CompactionConfig, TaskConfig, TeamConfig}, + memory::MemoryStore, + provider::{ContentBlockDelta, ContentBlockStart, ProviderEvent}, + runtime::{AgentStore, Runtime, RuntimePolicy, TaskStore, VolatileRuntimeStore}, +}; + +use super::support::{ScriptedProvider, StaticTool, model_info, ok_stream}; + +#[tokio::test] +async fn volatile_run_leaves_no_durable_trace_on_disk() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_path("volatile-notrace-tasks"); + let team_dir = temp_path("volatile-notrace-team"); + let transcript_dir = temp_path("volatile-notrace-transcripts"); + let store = VolatileRuntimeStore::new(); + + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + task_tool_stream("tool-1", "task_create", r#"{"subject":"write the report"}"#), + text_stream("report task created"), + ], + ); + + let runtime = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let config = volatile_config(tasks_dir.clone(), team_dir.clone(), transcript_dir.clone()); + let mut agent = runtime + .spawn_with_config("primary", model.clone(), config) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "start".to_string(), + }]) + .await + .expect("run completes"); + let agent_id = agent.id().to_string(); + + // Give the detached post-run memory-ingest task a chance to run before + // asserting on the filesystem — it must not create anything either. + tokio::time::sleep(Duration::from_millis(50)).await; + + assert!( + !tasks_dir.exists(), + "tasks_dir must never be created by the volatile profile" + ); + assert!( + !team_dir.exists(), + "team_dir must never be created by the volatile profile" + ); + assert!( + !transcript_dir.exists(), + "transcript_dir must never be created by the volatile profile" + ); + + // The run's effects are real, just in-memory: the tool call landed in + // the retained store, and ingest wrote the episode into it too. + assert_eq!( + store + .load_tasks(&tasks_dir) + .expect("load tasks") + .into_iter() + .map(|task| task.subject) + .collect::>(), + vec!["write the report".to_string()] + ); + assert!( + !store + .search_records(&agent_id, "report", 10) + .expect("search ingested memory") + .is_empty(), + "the detached memory-ingest task should have written into the volatile store" + ); +} + +#[tokio::test] +async fn volatile_store_truncates_without_creating_spill_artifacts() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_path("volatile-truncation-tasks"); + let team_dir = temp_path("volatile-truncation-team"); + let transcript_dir = temp_path("volatile-truncation-transcripts"); + let spill_dir = transcript_dir.join("tool-output"); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + task_tool_stream("tool-output", "oversized_output", r#"{}"#), + text_stream("done"), + ], + ); + + let runtime = Runtime::empty_builder() + .with_store(VolatileRuntimeStore::new()) + .with_provider_instance(provider) + .with_policy( + RuntimePolicy::default() + .with_max_tool_result_bytes(usize::MAX) + .with_max_tool_result_lines(1), + ) + .with_tool(StaticTool::success("oversized_output", "one\ntwo\nthree")) + .build() + .expect("build runtime"); + let config = volatile_config(tasks_dir.clone(), team_dir.clone(), transcript_dir.clone()); + let mut agent = runtime + .spawn_with_config("primary", model, config) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::Text { + text: "run the oversized tool".to_string(), + }]) + .await + .expect("run completes"); + + let content = match agent.history()[2].content.first().expect("tool result") { + ContentBlock::ToolResult { + content, is_error, .. + } => { + assert!(!is_error); + content.to_display_string() + } + other => panic!("unexpected content block: {other:?}"), + }; + let tasks_dir_exists = tasks_dir.exists(); + let team_dir_exists = team_dir.exists(); + let transcript_dir_exists = transcript_dir.exists(); + let spill_dir_exists = spill_dir.exists(); + + for path in [&tasks_dir, &team_dir, &transcript_dir] { + if path.exists() { + fs::remove_dir_all(path).expect("remove unexpected volatile artifact directory"); + } + } + + assert_eq!( + content, + "one\n[truncated: showing 1 of 3 lines; full output was not saved because the runtime store forbids durable artifacts]" + ); + assert!(!tasks_dir_exists, "volatile task artifacts must not exist"); + assert!(!team_dir_exists, "volatile team artifacts must not exist"); + assert!( + !transcript_dir_exists, + "volatile transcript artifacts must not exist" + ); + assert!(!spill_dir_exists, "volatile spill artifacts must not exist"); +} + +#[tokio::test] +async fn sequential_runs_on_retained_store_do_not_leak_records() { + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_path("volatile-isolation-tasks"); + let team_dir = temp_path("volatile-isolation-team"); + let transcript_dir = temp_path("volatile-isolation-transcripts"); + let store = VolatileRuntimeStore::new(); + let config = volatile_config(tasks_dir.clone(), team_dir.clone(), transcript_dir.clone()); + + // --- Run 1: same team_dir/tasks_dir/agent name as run 2 below, which is + // exactly the shared-default scenario the volatile profile's isolation + // contract has to defend against on a retained store. --- + let provider_1 = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + task_tool_stream("tool-1", "task_create", r#"{"subject":"first run task"}"#), + text_stream("first run complete"), + ], + ); + let runtime_1 = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider_1) + .build() + .expect("build runtime 1"); + let mut agent_1 = runtime_1 + .spawn_with_config("primary", model.clone(), config.clone()) + .expect("spawn agent 1"); + agent_1 + .send(vec![ContentBlock::Text { + text: "go".to_string(), + }]) + .await + .expect("run 1 completes"); + let agent_1_id = agent_1.id().to_string(); + tokio::time::sleep(Duration::from_millis(50)).await; + + assert_eq!( + store.list_agents().expect("list agents after run 1").len(), + 1 + ); + assert_eq!( + store + .load_tasks(&tasks_dir) + .expect("tasks after run 1") + .len(), + 1 + ); + + // Explicit isolation seam: reset the retained store between runs. + store.reset(); + + // --- Run 2: a fresh agent (fresh id, same name/dirs) must observe none + // of run 1's records through the same retained store instance. --- + let provider_2 = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream("second run complete")], + ); + let runtime_2 = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider_2) + .build() + .expect("build runtime 2"); + let mut agent_2 = runtime_2 + .spawn_with_config("primary", model.clone(), config) + .expect("spawn agent 2"); + agent_2 + .send(vec![ContentBlock::Text { + text: "go".to_string(), + }]) + .await + .expect("run 2 completes"); + let agent_2_id = agent_2.id().to_string(); + tokio::time::sleep(Duration::from_millis(50)).await; + + assert_ne!(agent_1_id, agent_2_id, "each spawn gets a fresh agent id"); + + let agents_after_run_2 = store.list_agents().expect("list agents after run 2"); + assert_eq!( + agents_after_run_2.len(), + 1, + "run 2 must not see run 1's agent record" + ); + assert_eq!(agents_after_run_2[0].record.id, agent_2_id); + + assert!( + store + .load_tasks(&tasks_dir) + .expect("tasks after run 2") + .is_empty(), + "run 2 must not see run 1's task, which was written under the same tasks_dir" + ); + + assert!( + store + .search_records(&agent_1_id, "first run", 10) + .expect("search for agent 1's memory by its own id") + .is_empty(), + "agent 1's own record disappeared with reset(), so nothing can match its id" + ); + assert!( + !store + .search_records(&agent_2_id, "second run", 10) + .expect("search for agent 2's memory") + .is_empty(), + "run 2's own ingested memory is still visible to itself" + ); +} + +#[tokio::test] +async fn retained_store_without_reset_shares_state_across_runs() { + // Companion to `sequential_runs_on_retained_store_do_not_leak_records`: + // demonstrates that `reset()` is doing real work by showing what happens + // without it. A retained `VolatileRuntimeStore` is a shared database — + // exactly like two runs pointed at the same `SqliteRuntimeStore` path — + // when the host does not call `reset()` between runs. + let model = model_info("model", BuiltinProvider::Anthropic); + let tasks_dir = temp_path("volatile-shared-tasks"); + let team_dir = temp_path("volatile-shared-team"); + let transcript_dir = temp_path("volatile-shared-transcripts"); + let store = VolatileRuntimeStore::new(); + let config = volatile_config(tasks_dir.clone(), team_dir.clone(), transcript_dir.clone()); + + let provider_1 = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + task_tool_stream("tool-1", "task_create", r#"{"subject":"first run task"}"#), + text_stream("first run complete"), + ], + ); + let runtime_1 = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider_1) + .build() + .expect("build runtime 1"); + let mut agent_1 = runtime_1 + .spawn_with_config("primary", model.clone(), config.clone()) + .expect("spawn agent 1"); + agent_1 + .send(vec![ContentBlock::Text { + text: "go".to_string(), + }]) + .await + .expect("run 1 completes"); + + // No reset() here. + + let provider_2 = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream("second run complete")], + ); + let runtime_2 = Runtime::builder() + .with_store(store.clone()) + .with_provider_instance(provider_2) + .build() + .expect("build runtime 2"); + let mut agent_2 = runtime_2 + .spawn_with_config("primary", model.clone(), config) + .expect("spawn agent 2"); + agent_2 + .send(vec![ContentBlock::Text { + text: "go".to_string(), + }]) + .await + .expect("run 2 completes"); + + assert_eq!( + store.list_agents().expect("list agents").len(), + 2, + "without reset(), both agents' records remain in the shared store" + ); + assert_eq!( + store + .load_tasks(&tasks_dir) + .expect("tasks after both runs") + .len(), + 1, + "without reset(), run 1's task is still visible under the shared tasks_dir" + ); +} + +fn volatile_config(tasks_dir: PathBuf, team_dir: PathBuf, transcript_dir: PathBuf) -> AgentConfig { + AgentConfig { + task: TaskConfig { + tasks_dir, + reminder_threshold: 3, + }, + team: TeamConfig { + team_dir, + ..Default::default() + }, + compaction: CompactionConfig { + transcript_dir, + ..Default::default() + }, + ..Default::default() + } +} + +fn task_tool_stream( + tool_id: &str, + tool_name: &str, + input_json: &str, +) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{tool_id}"), + model: "model".to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: tool_id.to_string(), + name: tool_name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn text_stream(text: &str) -> super::support::StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: "model".to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +/// Builds a unique path under the system temp directory *without* creating +/// it — the whole point of these tests is to assert the volatile profile +/// never creates it either. +fn temp_path(label: &str) -> PathBuf { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + std::env::temp_dir().join(format!("mentra-{label}-{timestamp}-{unique}")) +} diff --git a/vendor/mentra/src/agent/tests/steering.rs b/vendor/mentra/src/agent/tests/steering.rs new file mode 100644 index 0000000..dd15b98 --- /dev/null +++ b/vendor/mentra/src/agent/tests/steering.rs @@ -0,0 +1,672 @@ +use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, +}; + +use async_trait::async_trait; +use tokio::sync::mpsc; + +use crate::{ + BuiltinProvider, ContentBlock, Message, QueueMode, Role, RoundContext, RoundDecision, + RoundStrategy, Runtime, SteeringHandle, + provider::{ContentBlockDelta, ContentBlockStart, ProviderError, ProviderEvent}, + runtime::{ + CancellationToken, CommandOutput, CommandRequest, RunOptions, RuntimeExecutor, + RuntimePolicy, VolatileRuntimeStore, + }, +}; + +use super::support::{ + ScriptedProvider, StaticTool, controlled_stream, erroring_stream, model_info, text_stream, + tool_use_stream, +}; + +#[tokio::test] +async fn live_steer_is_visible_in_the_next_provider_request() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (first_stream, first_tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![first_stream, text_stream(&model.id, "final")], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let steering = agent.steering_handle(); + + let drive = async { + wait_for_request_count(&provider_handle, 1).await; + steering.steer(vec![ContentBlock::text("focus on the API contract")]); + send_text_response(&first_tx, &model.id, "draft"); + drop(first_tx); + }; + let (result, ()) = tokio::join!( + agent.run(vec![ContentBlock::text("start")], RunOptions::default()), + drive + ); + + assert_eq!(result.expect("run succeeds").text(), "final"); + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 2); + assert!(request_contains( + &requests[1].messages, + "focus on the API contract" + )); +} + +#[tokio::test] +async fn follow_up_waits_for_the_would_stop_boundary() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "probe", r#"{"value":"x"}"#), + text_stream(&model.id, "tool round complete"), + text_stream(&model.id, "follow-up complete"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe", "ok")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + agent.follow_up(vec![ContentBlock::text("now produce the appendix")]); + + let result = agent + .run(vec![ContentBlock::text("start")], RunOptions::default()) + .await + .expect("run succeeds"); + + assert_eq!(result.text(), "follow-up complete"); + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 3); + assert!(!request_contains( + &requests[1].messages, + "now produce the appendix" + )); + assert!(request_contains( + &requests[2].messages, + "now produce the appendix" + )); +} + +#[tokio::test] +async fn queue_modes_drain_one_or_all_entries_per_boundary() { + let one_at_a_time = run_with_queue_mode(QueueMode::OneAtATime).await; + assert_eq!(one_at_a_time.len(), 3); + assert!(request_contains(&one_at_a_time[1], "first steer")); + assert!(!request_contains(&one_at_a_time[1], "second steer")); + assert!(request_contains(&one_at_a_time[2], "second steer")); + + let all = run_with_queue_mode(QueueMode::All).await; + assert_eq!(all.len(), 2); + assert!(request_contains(&all[1], "first steer")); + assert!(request_contains(&all[1], "second steer")); +} + +#[test] +fn clear_methods_remove_only_their_pending_queue() { + assert_eq!(QueueMode::default(), QueueMode::OneAtATime); + + let steering = SteeringHandle::default(); + steering.steer(vec![ContentBlock::text("steer")]); + steering.follow_up(vec![ContentBlock::text("follow-up")]); + + steering.clear_steer(); + assert!(steering.has_pending(), "the follow-up remains queued"); + + steering.steer(vec![ContentBlock::text("replacement steer")]); + steering.clear_follow_up(); + assert!( + steering.has_pending(), + "the replacement steer remains queued" + ); + + steering.clear_steer(); + assert!(!steering.has_pending()); +} + +#[tokio::test] +async fn follow_up_all_mode_drains_every_entry_at_the_would_stop_boundary() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "draft"), + text_stream(&model.id, "final"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let steering = agent.steering_handle(); + steering.set_follow_up_mode(QueueMode::All); + steering.follow_up(vec![ContentBlock::text("first follow-up")]); + steering.follow_up(vec![ContentBlock::text("second follow-up")]); + + let result = agent + .send(vec![ContentBlock::text("start")]) + .await + .expect("run succeeds"); + + assert_eq!(result.text(), "final"); + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 2); + assert!(request_contains(&requests[1].messages, "first follow-up")); + assert!(request_contains(&requests[1].messages, "second follow-up")); + assert!(!steering.has_pending()); +} + +#[tokio::test] +async fn failed_run_requeues_steer_and_resume_reinjects_it() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first draft"), + erroring_stream( + Vec::new(), + ProviderError::MalformedStream("failed after steer".to_string()), + ), + text_stream(&model.id, "retry draft"), + text_stream(&model.id, "fixed"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let steering = agent.steering_handle(); + steering.steer(vec![ContentBlock::text("repair the draft")]); + + agent + .run(vec![ContentBlock::text("start")], RunOptions::default()) + .await + .expect_err("second request fails"); + assert!(steering.has_pending(), "failed run requeues the steer"); + + let result = agent.resume().await.expect("resume succeeds"); + assert_eq!(result.text(), "fixed"); + assert!(!steering.has_pending()); + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 4); + assert!(request_contains(&requests[1].messages, "repair the draft")); + assert!(request_contains(&requests[3].messages, "repair the draft")); +} + +#[tokio::test] +async fn failed_run_requeues_follow_up_and_resume_reinjects_it() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "first draft"), + erroring_stream( + Vec::new(), + ProviderError::MalformedStream("failed after follow-up".to_string()), + ), + text_stream(&model.id, "retry draft"), + text_stream(&model.id, "fixed"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let steering = agent.steering_handle(); + steering.follow_up(vec![ContentBlock::text("append the required evidence")]); + + agent + .run(vec![ContentBlock::text("start")], RunOptions::default()) + .await + .expect_err("second request fails"); + assert!(steering.has_pending(), "failed run requeues the follow-up"); + + let result = agent.resume().await.expect("resume succeeds"); + assert_eq!(result.text(), "fixed"); + assert!(!steering.has_pending()); + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 4); + assert!(!request_contains( + &requests[0].messages, + "append the required evidence" + )); + assert!(request_contains( + &requests[1].messages, + "append the required evidence" + )); + assert!(!request_contains( + &requests[2].messages, + "append the required evidence" + )); + assert!(request_contains( + &requests[3].messages, + "append the required evidence" + )); + assert_eq!( + agent + .history() + .iter() + .filter(|message| message.text().contains("append the required evidence")) + .count(), + 1 + ); +} + +#[tokio::test] +async fn finalization_error_requeues_steering_before_returning() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (first_stream, first_tx) = controlled_stream(); + let (second_stream, second_tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + first_stream, + second_stream, + text_stream(&model.id, "queued steer completed"), + ], + ); + let provider_handle = provider.clone(); + let store = VolatileRuntimeStore::new(); + let runtime = Runtime::empty_builder() + .with_store(store.clone()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let steering = agent.steering_handle(); + steering.steer(vec![ContentBlock::text("preserve this steer")]); + + let drive = async { + wait_for_request_count(&provider_handle, 1).await; + send_text_response(&first_tx, &model.id, "draft"); + drop(first_tx); + + wait_for_request_count(&provider_handle, 2).await; + store.fail_next_agent_record_save(); + send_text_response(&second_tx, &model.id, "final"); + drop(second_tx); + }; + let (result, ()) = tokio::join!( + agent.run(vec![ContentBlock::text("start")], RunOptions::default()), + drive + ); + + assert!( + result + .expect_err("finalization persistence fails") + .to_string() + .contains("injected agent-record persistence failure") + ); + assert!(steering.has_pending(), "finalization error requeues steer"); + + let recovered = agent + .run_queued(RunOptions::default()) + .await + .expect("queued steer remains runnable"); + assert_eq!(recovered.text(), "queued steer completed"); + assert!(!steering.has_pending()); +} + +#[tokio::test] +async fn steering_precedes_round_strategy_without_double_injection() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "draft"), + text_stream(&model.id, "steered"), + text_stream(&model.id, "strategized"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + agent.steer(vec![ContentBlock::text("queue correction")]); + let strategy = Arc::new(InjectOnceStrategy::default()); + + agent + .run( + vec![ContentBlock::text("start")], + RunOptions::default().with_round_strategy(strategy.clone()), + ) + .await + .expect("run succeeds"); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 3); + assert!(request_contains(&requests[1].messages, "queue correction")); + assert!(!request_contains( + &requests[1].messages, + "strategy correction" + )); + assert!(request_contains( + &requests[2].messages, + "strategy correction" + )); + assert_eq!(strategy.calls.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn next_provider_request_orders_steering_team_inbox_then_background() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (first_stream, first_tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![first_stream, text_stream(&model.id, "final")], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_executor(ImmediateExecutor) + .with_policy(RuntimePolicy::permissive()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let steering = agent.steering_handle(); + let runtime_handle = agent.runtime.clone(); + let agent_id = agent.id.clone(); + let agent_name = agent.name.clone(); + let team_dir = agent.config.team.team_dir.clone(); + let cwd = agent.config.workspace.base_dir.clone(); + + let drive = async { + wait_for_request_count(&provider_handle, 1).await; + steering.steer(vec![ContentBlock::text("ordering-steer-marker")]); + runtime_handle + .send_team_message( + &team_dir, + "host", + &agent_name, + "ordering-team-marker".to_string(), + ) + .expect("enqueue team message"); + runtime_handle + .start_background_task( + &agent_id, + "ordering-background-marker".to_string(), + None, + None, + cwd, + ) + .expect("start background task"); + while !runtime_handle.has_deliverable_background_notifications(&agent_id) { + tokio::task::yield_now().await; + } + + send_text_response(&first_tx, &model.id, "draft"); + drop(first_tx); + }; + let (result, ()) = tokio::join!( + agent.run(vec![ContentBlock::text("start")], RunOptions::default()), + drive + ); + + assert_eq!(result.expect("run succeeds").text(), "final"); + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 2); + let messages = requests[1].messages.as_ref(); + let steer = message_index(messages, "ordering-steer-marker"); + let team = message_index(messages, "ordering-team-marker"); + let background = message_index(messages, "ordering-background-marker"); + assert!( + steer < team && team < background, + "expected steering -> team inbox -> background, got {:?}", + messages.iter().map(Message::text).collect::>() + ); +} + +#[tokio::test] +async fn run_queued_consumes_idle_steer_as_the_user_turn() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "done")], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + agent.steer(vec![ContentBlock::text("queued while idle")]); + + let result = agent + .run_queued(RunOptions::default()) + .await + .expect("queued run succeeds"); + assert_eq!(result.text(), "done"); + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 1); + assert!(request_contains(&requests[0].messages, "queued while idle")); +} + +#[tokio::test] +async fn steering_is_isolated_between_agents_on_one_runtime() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "agent-b done"), + text_stream(&model.id, "agent-a draft"), + text_stream(&model.id, "agent-a done"), + ], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent_a = runtime.spawn("agent-a", model.clone()).expect("spawn a"); + let mut agent_b = runtime.spawn("agent-b", model).expect("spawn b"); + agent_a.steer(vec![ContentBlock::text("only agent a sees this")]); + + agent_b + .send(vec![ContentBlock::text("run b")]) + .await + .expect("run b"); + agent_a + .send(vec![ContentBlock::text("run a")]) + .await + .expect("run a"); + + let requests = provider_handle.recorded_requests().await; + assert_eq!(requests.len(), 3); + assert!(!request_contains( + &requests[0].messages, + "only agent a sees this" + )); + assert!(request_contains( + &requests[2].messages, + "only agent a sees this" + )); +} + +#[tokio::test] +async fn graceful_stop_does_not_consume_an_unrequestable_steer() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (first_stream, first_tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![first_stream, text_stream(&model.id, "must not run")], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model.clone()).expect("spawn agent"); + let steering = agent.steering_handle(); + steering.steer(vec![ContentBlock::text("keep for later")]); + let stop = CancellationToken::default(); + let stop_driver = stop.clone(); + + let drive = async { + wait_for_request_count(&provider_handle, 1).await; + stop_driver.cancel(); + send_text_response(&first_tx, &model.id, "finished before steer"); + drop(first_tx); + }; + let (result, ()) = tokio::join!( + agent.run( + vec![ContentBlock::text("start")], + RunOptions { + stop: Some(stop), + ..RunOptions::default() + } + ), + drive + ); + + assert_eq!( + result.expect("graceful stop succeeds").text(), + "finished before steer" + ); + assert!(steering.has_pending()); + assert_eq!(provider_handle.recorded_requests().await.len(), 1); +} + +async fn run_with_queue_mode(mode: QueueMode) -> Vec> { + let model = model_info("model", BuiltinProvider::Anthropic); + let scripts = match mode { + QueueMode::OneAtATime => vec![ + text_stream(&model.id, "first"), + text_stream(&model.id, "second"), + text_stream(&model.id, "third"), + ], + QueueMode::All => vec![ + text_stream(&model.id, "first"), + text_stream(&model.id, "second"), + ], + }; + let provider = ScriptedProvider::new(BuiltinProvider::Anthropic, vec![model.clone()], scripts); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let steering = agent.steering_handle(); + steering.set_steer_mode(mode); + steering.steer(vec![ContentBlock::text("first steer")]); + steering.steer(vec![ContentBlock::text("second steer")]); + + agent + .send(vec![ContentBlock::text("start")]) + .await + .expect("run"); + provider_handle + .recorded_requests() + .await + .into_iter() + .map(|request| request.messages.to_vec()) + .collect() +} + +#[derive(Default)] +struct InjectOnceStrategy { + calls: AtomicUsize, +} + +#[async_trait] +impl RoundStrategy for InjectOnceStrategy { + async fn on_round(&self, _ctx: RoundContext<'_>) -> RoundDecision { + if self.calls.fetch_add(1, Ordering::SeqCst) == 0 { + RoundDecision::inject(vec![ContentBlock::text("strategy correction")]) + } else { + RoundDecision::stop() + } + } +} + +struct ImmediateExecutor; + +#[async_trait] +impl RuntimeExecutor for ImmediateExecutor { + async fn run(&self, _request: CommandRequest) -> Result { + Ok(CommandOutput { + stdout: "ordering background output".to_string(), + stderr: String::new(), + success: true, + status_code: Some(0), + timed_out: false, + stdout_truncated: false, + stderr_truncated: false, + }) + } +} + +fn request_contains(messages: &[Message], needle: &str) -> bool { + messages + .iter() + .any(|message| message.text().contains(needle)) +} + +fn message_index(messages: &[Message], needle: &str) -> usize { + messages + .iter() + .position(|message| message.text().contains(needle)) + .unwrap_or_else(|| panic!("request did not contain {needle:?}")) +} + +async fn wait_for_request_count(provider: &ScriptedProvider, expected: usize) { + loop { + if provider.recorded_requests().await.len() >= expected { + return; + } + tokio::task::yield_now().await; + } +} + +fn send_text_response( + tx: &mpsc::UnboundedSender>, + model: &str, + text: &str, +) { + let events = [ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]; + for event in events { + tx.send(Ok(event)).expect("stream receiver remains alive"); + } +} diff --git a/vendor/mentra/src/agent/tests/support.rs b/vendor/mentra/src/agent/tests/support.rs new file mode 100644 index 0000000..f0626a7 --- /dev/null +++ b/vendor/mentra/src/agent/tests/support.rs @@ -0,0 +1,494 @@ +use std::{ + collections::VecDeque, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use async_trait::async_trait; +use serde_json::{Value, json}; +use tokio::{ + sync::{Mutex, mpsc}, + time::sleep, +}; + +use crate::{ + Role, + provider::{ + CompactionRequest, CompactionResponse, ContentBlockDelta, ContentBlockStart, ModelInfo, + Provider, ProviderCapabilities, ProviderDescriptor, ProviderError, ProviderEvent, + ProviderEventStream, ProviderId, Request, + }, + tool::{ + ParallelToolContext, ToolContext, ToolDefinition, ToolExecutionCategory, ToolExecutor, + ToolResult, ToolSpec, + }, +}; + +pub(super) enum StreamScript { + Buffered(Vec>), + Receiver(ProviderEventStream), +} + +#[derive(Clone)] +pub(super) struct ScriptedProvider { + kind: ProviderId, + models: Vec, + scripts: Arc>>, + requests: Arc>>>, + compact_scripts: Arc>>>, + capabilities: ProviderCapabilities, +} + +impl ScriptedProvider { + pub(super) fn new( + kind: impl Into, + models: Vec, + scripts: Vec, + ) -> Self { + Self { + kind: kind.into(), + models, + scripts: Arc::new(Mutex::new(VecDeque::from(scripts))), + requests: Arc::new(Mutex::new(Vec::>::new())), + compact_scripts: Arc::new(Mutex::new(VecDeque::new())), + capabilities: ProviderCapabilities::default(), + } + } + + pub(super) async fn recorded_requests(&self) -> Vec> { + self.requests.lock().await.clone() + } + + pub(super) async fn push_compact_response( + &self, + response: Result, + ) { + self.compact_scripts.lock().await.push_back(response); + } + + pub(super) fn with_capabilities(mut self, capabilities: ProviderCapabilities) -> Self { + self.capabilities = capabilities; + self + } +} + +#[async_trait] +impl Provider for ScriptedProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.kind.clone()) + } + + fn capabilities(&self) -> ProviderCapabilities { + self.capabilities + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(self.models.clone()) + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.requests.lock().await.push(request.into_owned()); + match self.scripts.lock().await.pop_front() { + Some(StreamScript::Buffered(items)) => { + let (tx, rx) = mpsc::unbounded_channel(); + for item in items { + tx.send(item) + .expect("test stream receiver dropped unexpectedly"); + } + Ok(rx) + } + Some(StreamScript::Receiver(receiver)) => Ok(receiver), + None => panic!("no scripted stream available"), + } + } + + async fn compact( + &self, + _request: CompactionRequest<'_>, + ) -> Result { + match self.compact_scripts.lock().await.pop_front() { + Some(result) => result, + None => Err(ProviderError::UnsupportedCapability( + "history_compaction".to_string(), + )), + } + } +} + +pub(super) fn model_info(id: &str, provider: impl Into) -> ModelInfo { + ModelInfo::new(id, provider) +} + +pub(super) fn ok_stream(events: Vec) -> StreamScript { + StreamScript::Buffered(events.into_iter().map(Ok).collect()) +} + +pub(super) fn command_input_json(command: &str) -> String { + json!({ "command": command }).to_string() +} + +pub(super) fn command_input_with_working_directory_json( + command: &str, + working_directory: &str, +) -> String { + json!({ + "command": command, + "workingDirectory": working_directory, + }) + .to_string() +} + +pub(super) fn shell_pwd_command() -> String { + #[cfg(unix)] + { + "pwd".to_string() + } + + #[cfg(windows)] + { + "cd".to_string() + } +} + +pub(super) fn background_success_command(output: &str, delay_ms: u64) -> String { + #[cfg(unix)] + { + format!( + "sleep {}; printf {}", + delay_seconds(delay_ms), + shell_single_quoted(output) + ) + } + + #[cfg(windows)] + { + let delay_seconds = (delay_ms / 1000).saturating_add(1); + format!( + "ping -n {delay_seconds} 127.0.0.1 >NUL & echo {output}", + output = cmd_echo_literal(output) + ) + } +} + +pub(super) fn background_failure_command(stderr: &str, exit_code: i32, delay_ms: u64) -> String { + #[cfg(unix)] + { + format!( + "sleep {}; printf {} >&2; exit {exit_code}", + delay_seconds(delay_ms), + shell_single_quoted(stderr) + ) + } + + #[cfg(windows)] + { + let delay_seconds = (delay_ms / 1000).saturating_add(1); + format!( + "ping -n {delay_seconds} 127.0.0.1 >NUL & echo {stderr} 1>&2 & exit /b {exit_code}", + stderr = cmd_echo_literal(stderr) + ) + } +} + +#[cfg(unix)] +fn delay_seconds(delay_ms: u64) -> String { + format!("{:.3}", delay_ms as f64 / 1000.0) +} + +#[cfg(unix)] +fn shell_single_quoted(value: &str) -> String { + format!("'{}'", value.replace('\'', r"'\''")) +} + +#[cfg(windows)] +fn cmd_echo_literal(value: &str) -> String { + value + .replace('^', "^^") + .replace('&', "^&") + .replace('|', "^|") + .replace('<', "^<") + .replace('>', "^>") +} + +pub(super) fn erroring_stream(events: Vec, error: ProviderError) -> StreamScript { + let mut items = events.into_iter().map(Ok).collect::>(); + items.push(Err(error)); + StreamScript::Buffered(items) +} + +pub(super) fn controlled_stream() -> ( + StreamScript, + mpsc::UnboundedSender>, +) { + let (tx, rx) = mpsc::unbounded_channel(); + (StreamScript::Receiver(rx), tx) +} + +pub(super) struct StaticTool { + name: &'static str, + result: ToolResult, + loading_policy: crate::tool::ToolLoadingPolicy, +} + +impl StaticTool { + pub(super) fn success(name: &'static str, output: &str) -> Self { + Self { + name, + result: Ok(output.to_string()), + loading_policy: crate::tool::ToolLoadingPolicy::Immediate, + } + } + + pub(super) fn failure(name: &'static str, error: &str) -> Self { + Self { + name, + result: Err(error.to_string()), + loading_policy: crate::tool::ToolLoadingPolicy::Immediate, + } + } + + pub(super) fn deferred_success(name: &'static str, output: &str) -> Self { + Self { + name, + result: Ok(output.to_string()), + loading_policy: crate::tool::ToolLoadingPolicy::Deferred, + } + } +} + +#[async_trait] +impl ToolDefinition for StaticTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder(self.name) + .description("test tool") + .input_schema(json!({ + "type": "object", + "properties": { + "value": { "type": "string" } + } + })) + .side_effect_level(crate::tool::ToolSideEffectLevel::None) + .durability(crate::tool::ToolDurability::ReplaySafe) + .loading_policy(self.loading_policy) + .build() + } +} + +#[async_trait] +impl ToolExecutor for StaticTool { + async fn execute_mut(&self, _ctx: ToolContext<'_>, _input: Value) -> ToolResult { + self.result.clone() + } +} + +/// A tool that trips a graceful-stop token when executed, then succeeds — used to +/// exercise [`RunOptions::stop`] firing at the round boundary *after* a real tool +/// round, so the gathered transcript is committed rather than rolled back. +pub(super) struct StopTrippingTool { + name: &'static str, + stop: crate::runtime::CancellationToken, +} + +impl StopTrippingTool { + pub(super) fn new(name: &'static str, stop: crate::runtime::CancellationToken) -> Self { + Self { name, stop } + } +} + +#[async_trait] +impl ToolDefinition for StopTrippingTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder(self.name) + .description("test tool that requests a graceful stop") + .input_schema(json!({ + "type": "object", + "properties": { "value": { "type": "string" } } + })) + .side_effect_level(crate::tool::ToolSideEffectLevel::None) + .durability(crate::tool::ToolDurability::ReplaySafe) + .loading_policy(crate::tool::ToolLoadingPolicy::Immediate) + .build() + } +} + +#[async_trait] +impl ToolExecutor for StopTrippingTool { + async fn execute_mut(&self, _ctx: ToolContext<'_>, _input: Value) -> ToolResult { + self.stop.cancel(); + Ok("stopped".to_string()) + } +} + +#[derive(Clone)] +pub(super) struct ProbeTool { + name: &'static str, + parallel: bool, + delay: Duration, + log: Arc>>, + active: Arc, + max_active: Arc, +} + +impl ProbeTool { + pub(super) fn new( + name: &'static str, + parallel: bool, + delay: Duration, + log: Arc>>, + active: Arc, + max_active: Arc, + ) -> Self { + Self { + name, + parallel, + delay, + log, + active, + max_active, + } + } + + async fn run(&self) -> ToolResult { + self.log.lock().await.push(format!("{}:start", self.name)); + let active = self.active.fetch_add(1, Ordering::SeqCst) + 1; + let _ = self + .max_active + .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |current| { + (active > current).then_some(active) + }); + sleep(self.delay).await; + self.active.fetch_sub(1, Ordering::SeqCst); + self.log.lock().await.push(format!("{}:end", self.name)); + Ok(format!("{} complete", self.name)) + } +} + +#[async_trait] +impl ToolDefinition for ProbeTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder(self.name) + .description("probe tool") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .side_effect_level(crate::tool::ToolSideEffectLevel::None) + .durability(crate::tool::ToolDurability::ReplaySafe) + .build() + } +} + +#[async_trait] +impl ToolExecutor for ProbeTool { + fn execution_category(&self, _input: &Value) -> ToolExecutionCategory { + if self.parallel { + ToolExecutionCategory::ReadOnlyParallel + } else { + ToolExecutionCategory::ExclusiveLocalMutation + } + } + + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + self.run().await + } +} + +/// Creates a buffered `StreamScript` representing a single text response. +pub(super) fn text_stream(model: &str, text: &str) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +/// Creates a buffered `StreamScript` representing a single tool-use response. +pub(super) fn tool_use_stream(model: &str, id: &str, name: &str, input_json: &str) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{id}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: id.to_string(), + name: name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +/// Builder for generating multi-turn scripted sessions for testing. +pub(super) struct SessionGenerator { + scripts: Vec, + response_size: usize, + model_id: String, +} + +impl SessionGenerator { + pub(super) fn new(model_id: &str) -> Self { + Self { + scripts: Vec::new(), + response_size: 50, + model_id: model_id.to_string(), + } + } + + pub(super) fn with_response_size(mut self, chars: usize) -> Self { + self.response_size = chars; + self + } + + pub(super) fn add_text_turns(mut self, n: usize) -> Self { + for i in 0..n { + let text = format!( + "Response {i}: {}", + "x".repeat(self.response_size.saturating_sub(15)) + ); + self.scripts.push(text_stream(&self.model_id, &text)); + } + self + } + + #[allow(dead_code)] + pub(super) fn add_tool_turns(mut self, n: usize, tool_name: &str) -> Self { + for i in 0..n { + self.scripts.push(tool_use_stream( + &self.model_id, + &format!("tool-{i}"), + tool_name, + &format!(r#"{{"index":{i}}}"#), + )); + } + // Final text response after all tool calls + self.scripts.push(text_stream(&self.model_id, "tools done")); + self + } + + pub(super) fn build(self) -> Vec { + self.scripts + } +} diff --git a/vendor/mentra/src/agent/tests/terminal_output.rs b/vendor/mentra/src/agent/tests/terminal_output.rs new file mode 100644 index 0000000..b0e0e85 --- /dev/null +++ b/vendor/mentra/src/agent/tests/terminal_output.rs @@ -0,0 +1,684 @@ +//! What a typed turn ([`Agent::run_to_output`]) may do on its way to the +//! answer: the shaping turn that holds one forced tool, the working turn that +//! keeps its whole toolset, and what each does when the terminal call never +//! comes. + +use std::{ + collections::VecDeque, + sync::{Arc, Mutex}, +}; + +use async_trait::async_trait; +use serde::Deserialize; +use serde_json::{Value, json}; + +use crate::{ + AgentConfig, BuiltinProvider, ContentBlock, ModelInfo, Provider, ProviderDescriptor, + ProviderError, ProviderEventStream, Request, Role, Runtime, TerminalOutputSpec, TokenUsage, + error::RuntimeError, + provider::{Response, ToolChoice}, + provider_event_stream_from_response, + runtime::{CancellationToken, EarlyEnd, RunOptions}, +}; + +use super::support::{StaticTool, StopTrippingTool}; + +/// The prefix every generated terminal tool's name carries, which is how the +/// scripted model below finds a tool whose name it cannot know in advance. +const TERMINAL_PREFIX: &str = "mentra_terminal_"; + +#[derive(Debug, Deserialize, PartialEq, Eq)] +struct Review { + verdict: String, +} + +/// One block of a scripted assistant response. +/// +/// [`Say::Answer`] is resolved against the request's tool list when the round +/// runs, so a test never has to know the per-call name `run_to_output` +/// generates. +#[derive(Clone)] +enum Say { + Text(&'static str), + Call { + id: &'static str, + tool: &'static str, + }, + Answer { + id: &'static str, + input: Value, + }, +} + +/// One scripted round: what the model says, and what it reports having spent. +#[derive(Clone)] +struct Round { + blocks: Vec, + usage: Option, +} + +impl Round { + fn new(blocks: Vec) -> Self { + Self { + blocks, + usage: None, + } + } + + fn spending(mut self, input_tokens: u64, output_tokens: u64) -> Self { + self.usage = Some(TokenUsage { + input_tokens: Some(input_tokens), + output_tokens: Some(output_tokens), + ..Default::default() + }); + self + } +} + +/// What one request put in front of the model — the two things a typed turn +/// changes about a round. +#[derive(Clone, Debug)] +struct Offer { + tools: Vec, + choice: Option, +} + +impl Offer { + fn terminal_tool(&self) -> Option<&String> { + self.tools + .iter() + .find(|name| name.starts_with(TERMINAL_PREFIX)) + } + + fn ordinary_tools(&self) -> Vec<&String> { + self.tools + .iter() + .filter(|name| !name.starts_with(TERMINAL_PREFIX)) + .collect() + } +} + +/// A model that plays one scripted [`Round`] per request and records what each +/// request offered it. +#[derive(Clone)] +struct ScriptedModel { + model: ModelInfo, + rounds: Arc>>, + offers: Arc>>, +} + +impl ScriptedModel { + fn new(rounds: Vec) -> Self { + Self { + model: ModelInfo::new("typed-turn-model", BuiltinProvider::Anthropic), + rounds: Arc::new(Mutex::new(VecDeque::from(rounds))), + offers: Arc::new(Mutex::new(Vec::new())), + } + } + + fn offers(&self) -> Vec { + self.offers.lock().expect("offers poisoned").clone() + } +} + +#[async_trait] +impl Provider for ScriptedModel { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.model.provider.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(vec![self.model.clone()]) + } + + async fn stream(&self, request: Request<'_>) -> Result { + let offer = Offer { + tools: request.tools.iter().map(|tool| tool.name.clone()).collect(), + choice: request.tool_choice.clone(), + }; + let terminal = offer.terminal_tool().cloned(); + let index = { + let mut offers = self.offers.lock().expect("offers poisoned"); + offers.push(offer); + offers.len() - 1 + }; + let round = self + .rounds + .lock() + .expect("rounds poisoned") + .pop_front() + .unwrap_or_else(|| panic!("the model was asked for an unscripted round {index}")); + + let mut content = Vec::new(); + for block in round.blocks { + content.push(match block { + Say::Text(text) => ContentBlock::text(text), + Say::Call { id, tool } => ContentBlock::ToolUse { + id: id.to_string(), + name: tool.to_string(), + input: json!({ "value": "please" }), + }, + Say::Answer { id, input } => ContentBlock::ToolUse { + id: id.to_string(), + name: terminal + .clone() + .expect("the terminal tool must be on a typed turn's request"), + input, + }, + }); + } + let calls_a_tool = content + .iter() + .any(|block| matches!(block, ContentBlock::ToolUse { .. })); + + Ok(provider_event_stream_from_response(Response { + id: format!("message-{index}"), + model: self.model.id.clone(), + role: Role::Assistant, + content, + stop_reason: calls_a_tool.then(|| "tool_use".to_string()), + usage: round.usage, + })) + } +} + +fn review_spec() -> TerminalOutputSpec { + TerminalOutputSpec::new( + "submit_review", + "Return the verdict you reached", + json!({ + "type": "object", + "properties": { "verdict": { "type": "string" } }, + "required": ["verdict"] + }), + ) +} + +fn hold() -> Value { + json!({ "verdict": "hold" }) +} + +/// An agent that forces one ordinary tool of its own, so a test can tell a +/// typed turn's choice apart from the default one every agent already sends. +fn forcing_probe() -> AgentConfig { + AgentConfig { + tool_choice: Some(ToolChoice::Tool { + name: "probe".to_string(), + }), + ..AgentConfig::default() + } +} + +/// The `(tool_use_id, text, is_error)` of every result on the message that +/// ended the turn, in the order the round committed them. +fn last_results(message: &crate::Message) -> Vec<(String, String, bool)> { + message + .content + .iter() + .filter_map(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => Some((tool_use_id.clone(), content.to_display_string(), *is_error)), + _ => None, + }) + .collect() +} + +#[tokio::test] +async fn a_shaping_turn_offers_only_the_terminal_tool_and_forces_it() { + // The default typed turn, unchanged: an agent with a perfectly usable + // ordinary tool is not offered it, because the turn exists to decide a + // shape and nothing else. + let provider = ScriptedModel::new(vec![Round::new(vec![Say::Answer { + id: "answer-1", + input: hold(), + }])]); + let handle = provider.clone(); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe", "read the file")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("reviewer", model).expect("spawn agent"); + + let output = agent + .run_to_output::( + vec![ContentBlock::text("shape what you have")], + RunOptions::default(), + review_spec(), + ) + .await + .expect("a shaping turn answers"); + + assert_eq!(output.value.verdict, "hold"); + let offers = handle.offers(); + assert_eq!(offers.len(), 1, "a shaping turn takes one round"); + let terminal = offers[0] + .terminal_tool() + .expect("the terminal tool is offered") + .clone(); + assert_eq!( + offers[0].tools, + vec![terminal.clone()], + "the terminal tool is the only tool on the request" + ); + assert_eq!( + offers[0].choice, + Some(ToolChoice::Tool { name: terminal }), + "and the model is told to call it" + ); +} + +#[tokio::test] +async fn a_working_turn_reaches_an_ordinary_tool_and_then_answers_through_the_terminal_one() { + // The opt-in: one turn that reads and then answers in the declared shape, + // where the shaping turn would have needed a turn for each. The agent is + // configured to force a tool of its own, so the `Auto` below is this + // turn's doing and not a default. + let provider = ScriptedModel::new(vec![ + Round::new(vec![Say::Call { + id: "probe-1", + tool: "probe", + }]), + Round::new(vec![Say::Answer { + id: "answer-1", + input: hold(), + }]), + ]); + let handle = provider.clone(); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe", "read the file")) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("reviewer", model, forcing_probe()) + .expect("spawn agent"); + + let output = agent + .run_to_output::( + vec![ContentBlock::text("read, then review")], + RunOptions::default(), + review_spec().with_tools(), + ) + .await + .expect("a working turn answers"); + + assert_eq!(output.value.verdict, "hold"); + let offers = handle.offers(); + assert_eq!(offers.len(), 2, "the turn worked a round, then answered"); + for (round, offer) in offers.iter().enumerate() { + assert!( + offer.terminal_tool().is_some(), + "round {round} can end the turn" + ); + assert_eq!( + offer.ordinary_tools(), + vec!["probe"], + "round {round} keeps the ordinary toolset" + ); + assert_eq!( + offer.choice, + Some(ToolChoice::Auto), + "round {round} forces nothing: a forced choice — the agent's own \ + included — precludes the working rounds that are the point" + ); + } + + // The tool really ran — the point of the mode is the reading, not the + // roster. + let read_it = agent.history().iter().any(|message| { + message.content.iter().any(|block| { + matches!(block, ContentBlock::ToolResult { tool_use_id, content, .. } + if tool_use_id == "probe-1" && content.to_display_string() == "read the file") + }) + }); + assert!(read_it, "the ordinary tool executed: {:?}", agent.history()); +} + +#[tokio::test] +async fn a_working_turn_that_settles_for_prose_reports_the_missing_terminal_call() { + // Nothing forces the ending, so a model can work and then simply talk. + // That is not an answer, and it must not be reported as one. + let provider = ScriptedModel::new(vec![ + Round::new(vec![Say::Call { + id: "probe-1", + tool: "probe", + }]), + Round::new(vec![Say::Text("looks fine to me")]), + ]); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe", "read the file")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("reviewer", model).expect("spawn agent"); + + let error = agent + .run_to_output::( + vec![ContentBlock::text("read, then review")], + RunOptions::default(), + review_spec().with_tools(), + ) + .await + .expect_err("prose is not a typed answer"); + + assert!( + error + .to_string() + .contains("without invoking the expected terminal tool"), + "got: {error}" + ); + assert!( + agent + .history() + .iter() + .any(|message| message.text().contains("looks fine to me")), + "the turn keeps what it gathered and said" + ); +} + +#[tokio::test] +async fn a_working_turn_stopped_at_a_round_boundary_fails_instead_of_answering_nothing() { + // A working turn can run many rounds, so a graceful stop can now land + // between them. It ends the turn exactly as it ends any other — at the + // boundary, transcript kept — and the typed caller is told the terminal + // call never came rather than handed a value nobody produced. + let stop = CancellationToken::default(); + let provider = ScriptedModel::new(vec![ + Round::new(vec![Say::Call { + id: "probe-1", + tool: "stop_probe", + }]), + Round::new(vec![Say::Answer { + id: "answer-1", + input: hold(), + }]), + ]); + let handle = provider.clone(); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StopTrippingTool::new("stop_probe", stop.clone())) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("reviewer", model).expect("spawn agent"); + + let options = RunOptions { + stop: Some(stop), + ..Default::default() + }; + let error = agent + .run_to_output::( + vec![ContentBlock::text("read, then review")], + options.clone(), + review_spec().with_tools(), + ) + .await + .expect_err("a turn stopped before the terminal call has no value"); + + assert!( + error + .to_string() + .contains("without invoking the expected terminal tool"), + "got: {error}" + ); + assert_eq!( + options.ended_early(), + Some(EarlyEnd::StopRequested), + "and the run says which bound ended it" + ); + assert_eq!( + handle.offers().len(), + 1, + "the stop was honored at the boundary: the answering round never ran" + ); + assert_eq!( + agent.history().len(), + 3, + "the gathered round stays committed, not rolled back" + ); +} + +#[tokio::test] +async fn a_working_turn_out_of_token_budget_fails_the_same_way() { + // The other graceful bound, at the same boundary, reported as itself. + let provider = ScriptedModel::new(vec![ + Round::new(vec![Say::Call { + id: "probe-1", + tool: "probe", + }]) + .spending(60, 40), + Round::new(vec![Say::Answer { + id: "answer-1", + input: hold(), + }]), + ]); + let handle = provider.clone(); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe", "read the file")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("reviewer", model).expect("spawn agent"); + + let options = RunOptions { + token_budget: Some(100), + ..Default::default() + }; + let error = agent + .run_to_output::( + vec![ContentBlock::text("read, then review")], + options.clone(), + review_spec().with_tools(), + ) + .await + .expect_err("a turn out of budget before the terminal call has no value"); + + assert!( + error + .to_string() + .contains("without invoking the expected terminal tool"), + "got: {error}" + ); + assert_eq!(options.ended_early(), Some(EarlyEnd::TokenBudget)); + assert_eq!( + handle.offers().len(), + 1, + "the budget halted the run before the answering round" + ); +} + +#[tokio::test] +async fn a_terminal_call_beside_other_calls_ends_the_round_and_skips_what_follows() { + // A working turn is the first typed turn where the model can put other + // calls in the round it answers from. The terminal tool terminates its + // round, so calls before it run and calls after it do not — each still + // getting an explicit result, never a silent drop. + let provider = ScriptedModel::new(vec![Round::new(vec![ + Say::Call { + id: "before-1", + tool: "probe", + }, + Say::Answer { + id: "answer-1", + input: hold(), + }, + Say::Call { + id: "after-1", + tool: "probe", + }, + ])]); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe", "read the file")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("reviewer", model).expect("spawn agent"); + + let output = agent + .run_to_output::( + vec![ContentBlock::text("read and review in one breath")], + RunOptions::default(), + review_spec().with_tools(), + ) + .await + .expect("the terminal call in the round is still the answer"); + + assert_eq!(output.value.verdict, "hold"); + let results = last_results(&output.message); + assert_eq!( + results + .iter() + .map(|(id, _, _)| id.as_str()) + .collect::>(), + vec!["before-1", "answer-1", "after-1"], + "every call in the round has exactly one result" + ); + assert_eq!( + (results[0].1.as_str(), results[0].2), + ("read the file", false), + "the call before the answer ran" + ); + assert!( + results[2].1.contains("not executed: run terminated by") && results[2].2, + "the call after the answer did not run, and says so: {:?}", + results[2] + ); +} + +#[tokio::test] +async fn two_terminal_calls_in_one_round_answer_with_the_first() { + // The same rule read from the other side: the second terminal call is + // simply a call scheduled after a terminating one, so the first is the + // answer and the second is reported as skipped. Deliberate, because a + // model that emits two shapes has not told anyone which it meant. + let provider = ScriptedModel::new(vec![Round::new(vec![ + Say::Answer { + id: "answer-1", + input: json!({ "verdict": "hold" }), + }, + Say::Answer { + id: "answer-2", + input: json!({ "verdict": "ship" }), + }, + ])]); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("reviewer", model).expect("spawn agent"); + + let output = agent + .run_to_output::( + vec![ContentBlock::text("review it")], + RunOptions::default(), + review_spec().with_tools(), + ) + .await + .expect("the first terminal call answers"); + + assert_eq!(output.value.verdict, "hold"); + let results = last_results(&output.message); + assert_eq!(results.len(), 2); + assert!( + results[1].1.contains("not executed: run terminated by") && results[1].2, + "the second answer was never executed: {:?}", + results[1] + ); +} + +#[tokio::test] +async fn a_working_turn_leaves_the_gate_shut_behind_it() { + // The gate is per-run: whatever the typed turn did to the roster and to + // the choice, the next ordinary turn on the same agent is back to its own. + let provider = ScriptedModel::new(vec![ + Round::new(vec![Say::Answer { + id: "answer-1", + input: hold(), + }]), + Round::new(vec![Say::Text("back to prose")]), + ]); + let handle = provider.clone(); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("probe", "read the file")) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("reviewer", model, forcing_probe()) + .expect("spawn agent"); + + agent + .run_to_output::( + vec![ContentBlock::text("review it")], + RunOptions::default(), + review_spec().with_tools(), + ) + .await + .expect("a working turn answers"); + let plain = agent + .send(vec![ContentBlock::text("and now just talk")]) + .await + .expect("an ordinary turn follows"); + + assert_eq!(plain.text(), "back to prose"); + let offers = handle.offers(); + assert_eq!(offers[1].ordinary_tools(), vec!["probe"]); + assert!( + offers[1].terminal_tool().is_none(), + "the generated tool is gone once its run is over: {:?}", + offers[1] + ); + assert_eq!( + offers[1].choice, + Some(ToolChoice::Tool { + name: "probe".to_string() + }), + "and the agent's own forced choice is back" + ); +} + +/// A run that ends with no assistant message and no terminal call is reported +/// as the missing terminal call, not as the empty assistant response +/// `Agent::run` sees. Kept as its own test because it is the one place where +/// the typed helper deliberately reinterprets an error from underneath it. +#[tokio::test] +async fn a_run_that_answers_nothing_at_all_still_names_the_missing_terminal_call() { + let provider = ScriptedModel::new(vec![Round::new(Vec::new())]); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("reviewer", model).expect("spawn agent"); + + let error = agent + .run_to_output::( + vec![ContentBlock::text("review it")], + RunOptions::default(), + review_spec(), + ) + .await + .expect_err("an empty response is not an answer"); + + assert!( + !matches!(error, RuntimeError::EmptyAssistantResponse), + "the typed caller asked about the terminal call, not about prose" + ); + assert!( + error + .to_string() + .contains("without invoking the expected terminal tool"), + "got: {error}" + ); +} diff --git a/vendor/mentra/src/agent/tests/tool_output.rs b/vendor/mentra/src/agent/tests/tool_output.rs new file mode 100644 index 0000000..883bb48 --- /dev/null +++ b/vendor/mentra/src/agent/tests/tool_output.rs @@ -0,0 +1,1487 @@ +//! Tests for the M2 structured `ToolOutput` seam (ADR-0001 §3): the bridge +//! from `ToolResult`, structured content + opaque `details`, and the two +//! layers of termination exclusivity. + +use std::{ + collections::BTreeMap, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use serde::Deserialize; +use serde_json::{Value, json}; +use tokio::{ + sync::{Mutex as TokioMutex, mpsc}, + time::{Duration, sleep}, +}; + +use crate::{ + AgentConfig, BuiltinProvider, ContentBlock, FileToolProfile, Role, TerminalOutputSpec, + agent::{CompactionConfig, ToolProfile, WorkspaceConfig}, + provider::{ContentBlockDelta, ContentBlockStart, ProviderError, ProviderEvent, ToolChoice}, + runtime::{RunOptions, Runtime, RuntimeError, RuntimeHookEvent, RuntimePolicy}, + tool::{ + ParallelToolContext, ToolContext, ToolDefinition, ToolDurability, ToolExecutionCategory, + ToolExecutor, ToolOutput, ToolResult, ToolResultContent, ToolSideEffectLevel, ToolSpec, + }, +}; + +use super::support::{ + ScriptedProvider, StaticTool, StreamScript, controlled_stream, model_info, ok_stream, +}; + +/// A tool that only implements the new structured, exclusive-lane entry +/// point directly (no `execute_mut` override) — proves a tool need not touch +/// the old `String` surface at all to opt into structured content, details, +/// or termination. +struct StructuredDetailsTool; + +#[async_trait] +impl ToolDefinition for StructuredDetailsTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("structured_details_tool") + .description("test tool: returns structured content plus opaque details") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .build() + } +} + +#[async_trait] +impl ToolExecutor for StructuredDetailsTool { + async fn execute_mut_output( + &self, + _ctx: ToolContext<'_>, + _input: Value, + ) -> Result { + Ok( + ToolOutput::structured(json!({ "answer": 42 })) + .with_details(json!({ "secret": "shh" })), + ) + } +} + +/// A pair of parallel-eligible tools that each attach distinct `details`, +/// used to prove a round with several tool calls maps each result's +/// metadata to its own `tool_use_id` rather than to the wrong call or a +/// single collapsed value (M3 test 4). +struct DetailsToolA; + +#[async_trait] +impl ToolDefinition for DetailsToolA { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("details_tool_a") + .description("test tool: returns details keyed to call A") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .build() + } +} + +#[async_trait] +impl ToolExecutor for DetailsToolA { + async fn execute_output( + &self, + _ctx: ParallelToolContext, + _input: Value, + ) -> Result { + Ok(ToolOutput::text("a-result").with_details(json!({ "who": "a" }))) + } +} + +struct DetailsToolB; + +#[async_trait] +impl ToolDefinition for DetailsToolB { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("details_tool_b") + .description("test tool: returns details keyed to call B") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .build() + } +} + +#[async_trait] +impl ToolExecutor for DetailsToolB { + async fn execute_output( + &self, + _ctx: ParallelToolContext, + _input: Value, + ) -> Result { + Ok(ToolOutput::text("b-result").with_details(json!({ "who": "b" }))) + } +} + +/// A tool that ends the run as the value of its own execution. +struct TerminatingTool; + +#[async_trait] +impl ToolDefinition for TerminatingTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("terminating_tool") + .description("test tool: ends the run via ToolOutput::terminate") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .build() + } +} + +#[async_trait] +impl ToolExecutor for TerminatingTool { + async fn execute_mut_output( + &self, + _ctx: ToolContext<'_>, + _input: Value, + ) -> Result { + Ok(ToolOutput::text("final answer").terminating()) + } +} + +/// A tool that overrides the new structured surface directly and returns a +/// tool-level failure — proves `Err(String)` behaves identically on the new +/// entry point as it does on the old one. +struct FailingOutputTool; + +#[async_trait] +impl ToolDefinition for FailingOutputTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("failing_output_tool") + .description("test tool: fails via the new structured surface") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .build() + } +} + +struct ParallelOutputTool { + name: &'static str, + output: &'static str, + delay: Duration, +} + +#[async_trait] +impl ToolDefinition for ParallelOutputTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder(self.name) + .description("test tool: returns a large parallel result") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .build() + } +} + +#[async_trait] +impl ToolExecutor for ParallelOutputTool { + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + sleep(self.delay).await; + Ok(self.output.to_string()) + } +} + +struct OversizedStructuredTerminatingTool; + +#[async_trait] +impl ToolDefinition for OversizedStructuredTerminatingTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("oversized_structured_terminal") + .description("test tool: spills structured output while preserving metadata") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .terminal() + .build() + } +} + +#[async_trait] +impl ToolExecutor for OversizedStructuredTerminatingTool { + async fn execute_mut_output( + &self, + _ctx: ToolContext<'_>, + _input: Value, + ) -> Result { + Ok( + ToolOutput::structured(json!({ "payload": "x".repeat(128) })) + .with_details(json!({ "private": 42 })) + .terminating(), + ) + } +} + +#[async_trait] +impl ToolExecutor for FailingOutputTool { + async fn execute_mut_output( + &self, + _ctx: ToolContext<'_>, + _input: Value, + ) -> Result { + Err("boom".to_string()) + } +} + +/// A timing probe (start/end log, like the shared `ProbeTool`) that declares +/// `ReadOnlyParallel` on its descriptor but is also marked `.terminal()` — +/// used to prove the scheduler coerces a terminal-marked tool to exclusive +/// scheduling regardless of its declared category. +struct TerminalParallelProbe { + name: &'static str, + log: Arc>>, +} + +#[async_trait] +impl ToolDefinition for TerminalParallelProbe { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder(self.name) + .description("test tool: declares ReadOnlyParallel but is terminal") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .terminal() + .build() + } +} + +#[async_trait] +impl ToolExecutor for TerminalParallelProbe { + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + self.log.lock().await.push(format!("{}:start", self.name)); + sleep(Duration::from_millis(15)).await; + self.log.lock().await.push(format!("{}:end", self.name)); + Ok(format!("{} complete", self.name)) + } +} + +/// A genuinely parallel-eligible tool (not `.terminal()`) that nonetheless +/// tries to request termination from the parallel lane — the RUNTIME defense +/// scenario: misuse independent of the static descriptor marker. +struct MisbehavingParallelTerminateTool; + +#[async_trait] +impl ToolDefinition for MisbehavingParallelTerminateTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("misbehaving_parallel_terminate") + .description("test tool: wrongly requests termination from a parallel execution") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .build() + } +} + +#[async_trait] +impl ToolExecutor for MisbehavingParallelTerminateTool { + async fn execute_output( + &self, + _ctx: ParallelToolContext, + _input: Value, + ) -> Result { + Ok(ToolOutput::text("i should not be able to stop the run").terminating()) + } +} + +fn tool_use_stream(model: &str, id: &str, name: &str, input_json: &str) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{id}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: id.to_string(), + name: name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input_json.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn multi_tool_use_stream(model: &str, calls: &[(&str, &str, &str)]) -> StreamScript { + let mut events = vec![ProviderEvent::MessageStarted { + id: "msg-multi-tool".to_string(), + model: model.to_string(), + role: Role::Assistant, + }]; + + for (index, (id, name, input_json)) in calls.iter().enumerate() { + events.push(ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::ToolUse { + id: (*id).to_string(), + name: (*name).to_string(), + }, + }); + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ToolUseInputJson((*input_json).to_string()), + }); + events.push(ProviderEvent::ContentBlockStopped { index }); + } + + events.push(ProviderEvent::MessageStopped); + ok_stream(events) +} + +fn text_stream(model: &str, text: &str) -> StreamScript { + ok_stream(vec![ + ProviderEvent::MessageStarted { + id: format!("msg-{text}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} + +fn tool_result_blocks(messages: &[crate::Message]) -> Vec { + messages + .iter() + .filter(|message| message.role == Role::User) + .flat_map(|message| message.content.iter().cloned()) + .filter(|block| matches!(block, ContentBlock::ToolResult { .. })) + .collect() +} + +#[derive(Clone, Default)] +struct RecordingHook { + events: Arc>>, +} + +impl crate::runtime::control::RuntimeHook for RecordingHook { + fn on_event( + &self, + _store: &dyn crate::runtime::AuditStore, + event: &RuntimeHookEvent, + ) -> Result<(), RuntimeError> { + self.events + .lock() + .expect("hook events poisoned") + .push(event.clone()); + Ok(()) + } +} + +// 1. A string tool compiles unchanged and produces Text through the bridge, +// byte-identical transcript output vs today. +#[tokio::test] +async fn string_tool_bridges_to_text_output_unchanged() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "echo_tool", r#"{}"#), + text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::success("echo_tool", "echoed")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run the echo tool")]) + .await + .expect("send"); + + let blocks = tool_result_blocks(agent.history()); + assert_eq!(blocks.len(), 1); + assert_eq!( + blocks[0], + ContentBlock::ToolResult { + tool_use_id: "call-1".to_string(), + content: ToolResultContent::Text("echoed".to_string()), + is_error: false, + } + ); +} + +// 7. Err(String) from the new surface behaves exactly like today (is_error +// block, model sees it, run continues) — exercised both through the bridge +// (StaticTool::failure, unchanged `execute_mut`) and directly on the new +// `execute_mut_output` surface. +#[tokio::test] +async fn err_string_behaves_identically_through_bridge_and_new_surface() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-1", "bridged_failure", r#"{}"#), + ("call-2", "failing_output_tool", r#"{}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(StaticTool::failure("bridged_failure", "old bridge error")) + .with_tool(FailingOutputTool) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let message = agent + .send(vec![ContentBlock::text("run the failing tools")]) + .await + .expect("send should still succeed: tool errors don't fail the run"); + + assert_eq!(message.text(), "done"); + let blocks = tool_result_blocks(agent.history()); + assert_eq!( + blocks[0], + ContentBlock::ToolResult { + tool_use_id: "call-1".to_string(), + content: ToolResultContent::Text("old bridge error".to_string()), + is_error: true, + } + ); + assert_eq!( + blocks[1], + ContentBlock::ToolResult { + tool_use_id: "call-2".to_string(), + content: ToolResultContent::Text("boom".to_string()), + is_error: true, + } + ); +} + +#[tokio::test] +async fn exclusive_success_and_error_outputs_truncate_with_identical_rules() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-ok", "large_success", r#"{}"#), + ("call-err", "large_error", r#"{}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy( + RuntimePolicy::default() + .with_max_tool_result_bytes(usize::MAX) + .with_max_tool_result_lines(1) + .spill_full_tool_output(false), + ) + .with_tool(StaticTool::success("large_success", "one\ntwo\nthree")) + .with_tool(StaticTool::failure("large_error", "bad\nworse\nworst")) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run both tools")]) + .await + .expect("send"); + + let blocks = tool_result_blocks(agent.history()); + assert_eq!(blocks.len(), 2); + assert!(matches!( + &blocks[0], + ContentBlock::ToolResult { content: ToolResultContent::Text(content), is_error: false, .. } + if content.starts_with("one\n[truncated: showing 1 of 3 lines;") + )); + assert!(matches!( + &blocks[1], + ContentBlock::ToolResult { content: ToolResultContent::Text(content), is_error: true, .. } + if content.starts_with("bad\n[truncated: showing 1 of 3 lines;") + )); +} + +#[tokio::test] +async fn parallel_outputs_truncate_per_result_and_preserve_call_order() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-a", "large_parallel_a", r#"{}"#), + ("call-b", "large_parallel_b", r#"{}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy( + RuntimePolicy::default() + .with_max_tool_result_bytes(usize::MAX) + .with_max_tool_result_lines(1) + .spill_full_tool_output(false), + ) + .with_tool(ParallelOutputTool { + name: "large_parallel_a", + output: "a-one\na-two", + delay: Duration::from_millis(25), + }) + .with_tool(ParallelOutputTool { + name: "large_parallel_b", + output: "b-one\nb-two", + delay: Duration::from_millis(1), + }) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run both parallel tools")]) + .await + .expect("send"); + + let blocks = tool_result_blocks(agent.history()); + assert!(matches!( + &blocks[0], + ContentBlock::ToolResult { tool_use_id, content: ToolResultContent::Text(content), .. } + if tool_use_id == "call-a" && content.starts_with("a-one\n[truncated:") + )); + assert!(matches!( + &blocks[1], + ContentBlock::ToolResult { tool_use_id, content: ToolResultContent::Text(content), .. } + if tool_use_id == "call-b" && content.starts_with("b-one\n[truncated:") + )); +} + +#[tokio::test] +async fn oversized_structured_output_spills_without_losing_details_or_termination() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![tool_use_stream( + &model.id, + "call-structured", + "oversized_structured_terminal", + r#"{}"#, + )], + ); + let spill_root = std::env::temp_dir().join(format!( + "mentra-structured-spill-{}-{}", + std::process::id(), + SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos() + )); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(RuntimePolicy::default().with_max_tool_result_bytes(32)) + .with_tool(OversizedStructuredTerminatingTool) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + transcript_dir: spill_root.clone(), + ..Default::default() + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + + let result = agent + .run( + vec![ContentBlock::text("run structured tool")], + RunOptions::default(), + ) + .await; + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + + let blocks = tool_result_blocks(agent.history()); + assert!(matches!( + &blocks[0], + ContentBlock::ToolResult { content: ToolResultContent::Text(content), is_error: false, .. } + if content.contains("structured tool output") && content.contains("full output at") + )); + let item = agent + .transcript() + .items() + .iter() + .rev() + .find(|item| matches!(item.kind, crate::TranscriptKind::ToolExchange { .. })) + .expect("tool exchange item"); + assert_eq!( + item.detail("call-structured"), + Some(&json!({ "private": 42 })) + ); + + let output_dir = spill_root.join("tool-output"); + let files = std::fs::read_dir(&output_dir) + .expect("read spill directory") + .collect::, _>>() + .expect("read spill files"); + assert_eq!(files.len(), 1); + let stored = std::fs::read_to_string(files[0].path()).expect("read spill output"); + assert_eq!( + stored, + serde_json::to_string(&json!({ "payload": "x".repeat(128) })).unwrap() + ); + std::fs::remove_dir_all(spill_root).expect("remove spill root"); +} + +#[tokio::test] +async fn spill_failure_preserves_structured_details_and_termination() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![tool_use_stream( + &model.id, + "call-structured", + "oversized_structured_terminal", + r#"{}"#, + )], + ); + let blocking_file = std::env::temp_dir().join(format!( + "mentra-structured-spill-blocker-{}-{}", + std::process::id(), + SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos() + )); + std::fs::write(&blocking_file, "not a directory").expect("create blocking file"); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(RuntimePolicy::default().with_max_tool_result_bytes(32)) + .with_tool(OversizedStructuredTerminatingTool) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + compaction: CompactionConfig { + transcript_dir: blocking_file.clone(), + ..Default::default() + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + + let result = agent + .run( + vec![ContentBlock::text("run structured tool")], + RunOptions::default(), + ) + .await; + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + + let blocks = tool_result_blocks(agent.history()); + assert!(matches!( + &blocks[0], + ContentBlock::ToolResult { content: ToolResultContent::Text(content), is_error: false, .. } + if content.contains("full output could not be saved") + )); + let item = agent + .transcript() + .items() + .iter() + .rev() + .find(|item| matches!(item.kind, crate::TranscriptKind::ToolExchange { .. })) + .expect("tool exchange item"); + assert_eq!( + item.detail("call-structured"), + Some(&json!({ "private": 42 })) + ); + assert!(blocking_file.is_file()); + std::fs::remove_file(blocking_file).expect("remove blocking file"); +} + +// 2. A structured tool returns Structured content plus opaque details; the +// recorded provider request contains only the content projection, no +// details bytes. +#[tokio::test] +async fn structured_tool_projects_content_and_hides_details_from_provider() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "structured_details_tool", r#"{}"#), + text_stream(&model.id, "done"), + ], + ); + let hook = RecordingHook::default(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider.clone()) + .with_tool(StructuredDetailsTool) + .with_hook(hook.clone()) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run the structured tool")]) + .await + .expect("send"); + + let blocks = tool_result_blocks(agent.history()); + assert_eq!( + blocks[0], + ContentBlock::ToolResult { + tool_use_id: "call-1".to_string(), + content: ToolResultContent::Structured(json!({ "answer": 42 })), + is_error: false, + } + ); + + // The follow-up request (the one carrying the tool result back to the + // model) must contain the projected content and nothing from `details`. + let requests = provider.recorded_requests().await; + assert_eq!(requests.len(), 2, "tool round, then the follow-up round"); + let follow_up = requests[1] + .messages + .iter() + .map(|message| format!("{message:?}")) + .collect::>() + .join("\n"); + assert!(follow_up.contains("answer")); + assert!(!follow_up.contains("secret")); + assert!(!follow_up.contains("shh")); + + // `details` is still observable at the execution-outcome boundary via + // the runtime hook, for a host (or a later slice) to recover. + let events = hook.events.lock().expect("hook events poisoned").clone(); + let details = events.iter().find_map(|event| match event { + RuntimeHookEvent::ToolExecutionFinished { + tool_name, details, .. + } if tool_name == "structured_details_tool" => Some(details.clone()), + _ => None, + }); + assert_eq!(details, Some(Some(json!({ "secret": "shh" })))); +} + +#[tokio::test] +async fn edit_tool_details_never_enter_provider_projected_result_content() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream( + &model.id, + "call-edit", + "edit", + r#"{"path":"note.txt","edits":[{"old_string":"before","new_string":"after"}]}"#, + ), + text_stream(&model.id, "done"), + ], + ); + let unique = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("duration") + .as_nanos(); + let workspace = std::env::temp_dir().join(format!("mentra-edit-projection-{unique}")); + std::fs::create_dir_all(&workspace).expect("create workspace"); + std::fs::write(workspace.join("note.txt"), "before\n").expect("write note"); + let runtime = Runtime::builder() + .with_file_tools(FileToolProfile::Split) + .with_provider_instance(provider.clone()) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + workspace: WorkspaceConfig { + base_dir: workspace.clone(), + ..Default::default() + }, + ..Default::default() + }, + ) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("apply the requested edit")]) + .await + .expect("send"); + + let detail = agent + .transcript() + .items() + .iter() + .find_map(|item| item.detail("call-edit")) + .expect("edit details must survive locally"); + assert!(detail.get("diff").is_some()); + assert!(detail.get("patch").is_some()); + assert_eq!(detail["first_changed_line"], json!(1)); + + let requests = provider.recorded_requests().await; + let projected = requests[1] + .messages + .iter() + .flat_map(|message| message.content.iter()) + .find_map(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + content, + .. + } if tool_use_id == "call-edit" => Some(content), + _ => None, + }) + .expect("provider-projected edit result"); + assert_eq!( + projected, + &ToolResultContent::Text("Replaced 1 block in note.txt".to_string()) + ); + + std::fs::remove_dir_all(workspace).expect("remove workspace"); +} + +// M3 tests 3 & 4: a round with two parallel tool calls, each returning +// distinct `details`, maps every result's metadata to its own `tool_use_id` +// on the committed transcript item — not to the wrong call, and not +// collapsed into one value — and a host recovers it afterward through the +// plain public `TranscriptItem` accessor, with no mentra-side host type or +// downcast involved. +#[tokio::test] +async fn parallel_round_maps_each_results_details_to_its_own_tool_use_id() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-a", "details_tool_a", r#"{}"#), + ("call-b", "details_tool_b", r#"{}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(DetailsToolA) + .with_tool(DetailsToolB) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run both details tools")]) + .await + .expect("send"); + + // Public accessor, reached straight off `Agent::transcript()` — no + // downcast, no mentra-side knowledge of what "who" means. + let item = agent + .transcript() + .items() + .iter() + .rev() + .find(|item| matches!(item.kind, crate::TranscriptKind::ToolExchange { .. })) + .expect("tool exchange item committed"); + + let expected: BTreeMap = BTreeMap::from([ + ("call-a".to_string(), json!({ "who": "a" })), + ("call-b".to_string(), json!({ "who": "b" })), + ]); + assert_eq!(item.details(), Some(&expected)); + assert_eq!(item.detail("call-a"), Some(&json!({ "who": "a" }))); + assert_eq!(item.detail("call-b"), Some(&json!({ "who": "b" }))); +} + +// 3. A terminate:true tool ends the run successfully with the transcript +// committed. Mirrors `run_options_stop_after_tool_round_commits_transcript_and_halts`: +// the outer `Agent::run`/`send` returns `Err(EmptyAssistantResponse)` because +// the last committed message is the tool result (not assistant text) — this +// is the established, documented "honest stop, not a failure" contract for +// any tool-driven round ending (see that test and `team/actor.rs`'s +// `Ok(_) | Err(EmptyAssistantResponse) => Ok(())` handling); the run itself +// is NOT rolled back. +#[tokio::test] +async fn terminating_tool_commits_transcript_without_rollback() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + tool_use_stream(&model.id, "call-1", "terminating_tool", r#"{}"#), + // Must NOT be consumed: termination halts the run before this + // round's model request is issued. + text_stream(&model.id, "must not run"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(TerminatingTool) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let result = agent + .run( + vec![ContentBlock::text("go")], + RunOptions { + ..Default::default() + }, + ) + .await; + + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + assert_eq!( + agent.history().len(), + 3, + "user, the assistant tool call, and the committed tool result" + ); + let blocks = tool_result_blocks(agent.history()); + assert_eq!( + blocks[0], + ContentBlock::ToolResult { + tool_use_id: "call-1".to_string(), + content: ToolResultContent::Text("final answer".to_string()), + is_error: false, + } + ); +} + +// 4. A terminating call is never scheduled concurrently with retrieval +// (barrier test mirroring `parallel_batches_respect_exclusive_barriers`), +// and 5. calls scheduled after it in the same round get explicit +// not-executed results, in call order — exercised together since the +// skipped batch here is itself a parallel batch. +#[tokio::test] +async fn terminating_call_creates_a_barrier_and_skips_the_scheduled_parallel_batch() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![multi_tool_use_stream( + &model.id, + &[ + ("call-1", "probe_one", r#"{}"#), + ("call-2", "probe_two", r#"{}"#), + ("call-3", "terminating_tool", r#"{}"#), + ("call-4", "probe_three", r#"{}"#), + ], + )], + ); + let log = Arc::new(TokioMutex::new(Vec::new())); + let active = Arc::new(AtomicUsize::new(0)); + let max_active = Arc::new(AtomicUsize::new(0)); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(super::support::ProbeTool::new( + "probe_one", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(super::support::ProbeTool::new( + "probe_two", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(TerminatingTool) + .with_tool(super::support::ProbeTool::new( + "probe_three", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let result = agent + .run(vec![ContentBlock::text("go")], RunOptions::default()) + .await; + assert!(matches!(result, Err(RuntimeError::EmptyAssistantResponse))); + + // probe_three's batch was scheduled after the terminating call and must + // never have run. + let log = log.lock().await.clone(); + assert!(!log.contains(&"probe_three:start".to_string())); + assert!(!log.contains(&"probe_three:end".to_string())); + assert!( + log.contains(&"probe_one:start".to_string()) + && log.contains(&"probe_two:start".to_string()), + "the earlier parallel batch still ran before the barrier" + ); + assert!(max_active.load(Ordering::SeqCst) >= 2); + + // call-4 (probe_three) gets an explicit not-executed error result, after + // call-3's real terminate result, in call order. + let blocks = tool_result_blocks(agent.history()); + assert_eq!(blocks.len(), 4); + assert_eq!( + blocks[2], + ContentBlock::ToolResult { + tool_use_id: "call-3".to_string(), + content: ToolResultContent::Text("final answer".to_string()), + is_error: false, + } + ); + let ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } = &blocks[3] + else { + panic!("expected a tool result block"); + }; + assert_eq!(tool_use_id, "call-4"); + assert!(*is_error); + let text = content.to_display_string(); + assert!(text.contains("not executed")); + assert!(text.contains("terminating_tool")); +} + +// 6a. A terminal-marked tool declared with a parallel category is coerced +// to exclusive scheduling (STATIC layer). +#[tokio::test] +async fn terminal_marked_tool_declared_parallel_is_coerced_to_exclusive() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-1", "probe_one", r#"{}"#), + ("call-2", "probe_two", r#"{}"#), + ("call-3", "terminal_parallel_probe", r#"{}"#), + ("call-4", "probe_three", r#"{}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let log = Arc::new(TokioMutex::new(Vec::new())); + let active = Arc::new(AtomicUsize::new(0)); + let max_active = Arc::new(AtomicUsize::new(0)); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(super::support::ProbeTool::new( + "probe_one", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(super::support::ProbeTool::new( + "probe_two", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .with_tool(TerminalParallelProbe { + name: "terminal_parallel_probe", + log: Arc::clone(&log), + }) + .with_tool(super::support::ProbeTool::new( + "probe_three", + true, + Duration::from_millis(30), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text( + "run probes with a terminal barrier", + )]) + .await + .expect("send"); + + let log = log.lock().await.clone(); + let position = |entry: &str| { + log.iter() + .position(|logged| logged == entry) + .unwrap_or_else(|| panic!("missing log entry: {entry}")) + }; + let terminal_start = position("terminal_parallel_probe:start"); + let terminal_end = position("terminal_parallel_probe:end"); + let probe_one_end = position("probe_one:end"); + let probe_two_end = position("probe_two:end"); + let probe_three_start = position("probe_three:start"); + + // Deliberately compares END markers, not just start order: if the + // terminal marker were ignored, all four probes would share a single + // parallel batch (all declare ReadOnlyParallel) and run concurrently — + // the terminal probe's shorter 15ms delay would then very likely make it + // finish (`terminal_end`) *before* the 30ms probes even start their own + // end, so `probe_one_end < terminal_start` would fail. Only genuine + // exclusive-batch sequencing (this batch fully awaited before the next + // begins) guarantees the terminal probe starts after the earlier batch + // is entirely done, and the later batch starts after the terminal probe + // is entirely done. + assert!( + probe_one_end < terminal_start, + "terminal probe must not start until probe_one has fully finished" + ); + assert!( + probe_two_end < terminal_start, + "terminal probe must not start until probe_two has fully finished" + ); + assert!( + terminal_end < probe_three_start, + "probe_three must not start until the terminal probe has fully finished" + ); + assert!( + max_active.load(Ordering::SeqCst) >= 2, + "the surrounding probes still ran in parallel with each other" + ); +} + +// 6b. A parallel-lane terminate is rejected as an error result, not honored +// (RUNTIME layer) — independent of the static marker: this tool never +// declares `.terminal()`, it just misbehaves at runtime. +#[tokio::test] +async fn parallel_lane_terminate_is_rejected_as_misuse_and_run_continues() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-1", "misbehaving_parallel_terminate", r#"{}"#), + ("call-2", "probe_one", r#"{}"#), + ], + ), + text_stream(&model.id, "done"), + ], + ); + let log = Arc::new(TokioMutex::new(Vec::new())); + let active = Arc::new(AtomicUsize::new(0)); + let max_active = Arc::new(AtomicUsize::new(0)); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(MisbehavingParallelTerminateTool) + .with_tool(super::support::ProbeTool::new( + "probe_one", + true, + Duration::from_millis(10), + Arc::clone(&log), + Arc::clone(&active), + Arc::clone(&max_active), + )) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let message = agent + .send(vec![ContentBlock::text("run the misbehaving tool")]) + .await + .expect("send should succeed: the misuse is a tool error, not a run failure"); + + // The run proceeded to the follow-up round instead of ending — proof the + // bogus terminate was never honored. + assert_eq!(message.text(), "done"); + + let blocks = tool_result_blocks(agent.history()); + let ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } = &blocks[0] + else { + panic!("expected a tool result block"); + }; + assert_eq!(tool_use_id, "call-1"); + assert!(*is_error); + let text = content.to_display_string(); + assert!(text.contains("not honored")); + assert!(text.contains("parallel")); +} + +#[derive(Debug, Deserialize, PartialEq, Eq)] +struct TypedAnswer { + answer: u64, + evidence: Vec, +} + +#[tokio::test] +async fn run_to_output_forces_scoped_terminal_tool_and_extracts_exact_call_detail() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (stream, tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![stream], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config( + "target", + model.clone(), + AgentConfig { + tool_profile: ToolProfile::only(["ordinary_tool_only"]), + ..AgentConfig::default() + }, + ) + .expect("spawn target"); + let other_agent = runtime.spawn("other", model).expect("spawn other"); + let steering = agent.steering_handle(); + steering.steer(vec![ContentBlock::text("keep this queued")]); + + let drive = async { + wait_for_recorded_requests(&provider_handle, 1).await; + let requests = provider_handle.recorded_requests().await; + let ToolChoice::Tool { name } = requests[0] + .tool_choice + .clone() + .expect("terminal tool choice") + else { + panic!("run_to_output must force its terminal tool"); + }; + assert_eq!(requests[0].tools.len(), 1); + assert_eq!(requests[0].tools[0].name, name); + assert!(name.starts_with("mentra_terminal_finish_report_")); + assert!( + other_agent.tools().iter().all(|tool| tool.name != name), + "a scoped terminal tool must not leak into another agent's request" + ); + send_tool_response( + &tx, + "model", + "terminal-call-42", + &name, + r#"{"answer":42,"evidence":["a","b"]}"#, + ); + drop(tx); + name + }; + let (result, tool_name) = tokio::join!( + agent.run_to_output::( + vec![ContentBlock::text("produce typed output")], + RunOptions::default(), + TerminalOutputSpec::new( + "finish-report", + "Return the final report", + json!({ + "type": "object", + "properties": { + "answer": { "type": "integer" }, + "evidence": { "type": "array", "items": { "type": "string" } } + }, + "required": ["answer", "evidence"] + }), + ), + ), + drive + ); + + let output = result.expect("typed terminal output succeeds"); + assert_eq!( + output.value, + TypedAnswer { + answer: 42, + evidence: vec!["a".to_string(), "b".to_string()], + } + ); + assert_eq!(output.message.role, Role::User); + assert!(matches!( + output.message.content.as_slice(), + [ContentBlock::ToolResult { tool_use_id, .. }] if tool_use_id == "terminal-call-42" + )); + let last = agent.transcript().items().last().expect("terminal item"); + assert_eq!( + last.detail("terminal-call-42"), + Some(&json!({ "answer": 42, "evidence": ["a", "b"] })) + ); + assert!( + steering.has_pending(), + "terminal end_turn precedes steering" + ); + assert!( + runtime.tool_descriptor(&tool_name).is_none(), + "the generated tool is unregistered after the run" + ); +} + +#[tokio::test] +async fn run_to_output_never_reuses_stale_terminal_details() { + let model = model_info("model", BuiltinProvider::Anthropic); + let (first_stream, first_tx) = controlled_stream(); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![first_stream, text_stream(&model.id, "plain answer")], + ); + let provider_handle = provider.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let drive = async { + wait_for_recorded_requests(&provider_handle, 1).await; + let requests = provider_handle.recorded_requests().await; + let ToolChoice::Tool { name } = requests[0] + .tool_choice + .clone() + .expect("terminal tool choice") + else { + panic!("expected forced tool"); + }; + send_tool_response( + &first_tx, + "model", + "first-terminal-call", + &name, + r#"{"answer":1,"evidence":[]}"#, + ); + drop(first_tx); + }; + let (first, ()) = tokio::join!( + agent.run_to_output::( + vec![ContentBlock::text("first")], + RunOptions::default(), + terminal_spec(), + ), + drive + ); + assert_eq!(first.expect("first output").value.answer, 1); + + let error = agent + .run_to_output::( + vec![ContentBlock::text("second")], + RunOptions::default(), + terminal_spec(), + ) + .await + .expect_err("plain response must not reuse prior details"); + assert!( + error + .to_string() + .contains("without invoking the expected terminal tool") + ); +} + +fn terminal_spec() -> TerminalOutputSpec { + TerminalOutputSpec::new( + "finish", + "Return typed output", + json!({ + "type": "object", + "properties": { + "answer": { "type": "integer" }, + "evidence": { "type": "array", "items": { "type": "string" } } + }, + "required": ["answer", "evidence"] + }), + ) +} + +async fn wait_for_recorded_requests(provider: &ScriptedProvider, expected: usize) { + loop { + if provider.recorded_requests().await.len() >= expected { + return; + } + tokio::task::yield_now().await; + } +} + +fn send_tool_response( + tx: &mpsc::UnboundedSender>, + model: &str, + id: &str, + name: &str, + input: &str, +) { + let events = [ + ProviderEvent::MessageStarted { + id: format!("message-{id}"), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: id.to_string(), + name: name.to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson(input.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]; + for event in events { + tx.send(Ok(event)).expect("stream receiver remains alive"); + } +} diff --git a/vendor/mentra/src/agent/tests/tool_paging.rs b/vendor/mentra/src/agent/tests/tool_paging.rs new file mode 100644 index 0000000..92636de --- /dev/null +++ b/vendor/mentra/src/agent/tests/tool_paging.rs @@ -0,0 +1,566 @@ +//! End-to-end coverage for automatic tool-result paging: what the model sees, +//! what the event stream keeps, and how `read_tool_result` walks a retained +//! result window by window. + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::{ + AgentConfig, BuiltinProvider, ContentBlock, Message, Role, ToolResultPagingConfig, + agent::AgentEvent, + provider::{ContentBlockDelta, ContentBlockStart, ProviderEvent}, + runtime::{Runtime, RuntimePolicy}, + tool::{ + ParallelToolContext, ToolDefinition, ToolDurability, ToolExecutionCategory, ToolExecutor, + ToolResult, ToolSideEffectLevel, ToolSpec, + }, +}; + +use super::support::{ScriptedProvider, StaticTool, StreamScript, model_info, ok_stream}; + +/// Builds `count` lines of exactly 20 bytes each (`{tag}-{n:03}` padded), so +/// every window boundary asserted below is exact arithmetic on line counts +/// rather than an approximation. +fn numbered_lines(tag: &str, count: usize) -> String { + assert_eq!( + tag.len(), + 4, + "the fixed 20-byte line layout assumes a 4-byte tag" + ); + (1..=count) + .map(|line| format!("{tag}-{line:03}{}\n", "x".repeat(11))) + .collect() +} + +/// Removes the runtime's own tool-result caps from the picture. Paging runs +/// downstream of that limiter, so a threshold above `max_tool_result_bytes` +/// would never be reached — every paging test has to raise the caps first, +/// exactly as a real consumer enabling paging must. +fn unlimited_results() -> RuntimePolicy { + RuntimePolicy::default() + .with_max_tool_result_bytes(usize::MAX) + .with_max_tool_result_lines(usize::MAX) + .spill_full_tool_output(false) +} + +fn paged_config(threshold_bytes: usize, page_bytes: usize) -> AgentConfig { + AgentConfig { + tool_result_paging: Some(ToolResultPagingConfig { + threshold_bytes, + page_bytes, + }), + ..Default::default() + } +} + +fn tool_results(messages: &[Message]) -> Vec<(String, String, bool)> { + messages + .iter() + .filter(|message| message.role == Role::User) + .flat_map(|message| message.content.iter()) + .filter_map(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => Some((tool_use_id.clone(), content.to_display_string(), *is_error)), + _ => None, + }) + .collect() +} + +fn multi_tool_use_stream(model: &str, calls: &[(&str, &str, &str)]) -> StreamScript { + let mut events = vec![ProviderEvent::MessageStarted { + id: "msg-multi-tool".to_string(), + model: model.to_string(), + role: Role::Assistant, + }]; + for (index, (id, name, input_json)) in calls.iter().enumerate() { + events.push(ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::ToolUse { + id: (*id).to_string(), + name: (*name).to_string(), + }, + }); + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ToolUseInputJson((*input_json).to_string()), + }); + events.push(ProviderEvent::ContentBlockStopped { index }); + } + events.push(ProviderEvent::MessageStopped); + ok_stream(events) +} + +fn read_window_input(tool_use_id: &str, start_line: usize) -> String { + json!({ "tool_use_id": tool_use_id, "start_line": start_line }).to_string() +} + +/// A parallel-lane tool returning a caller-supplied oversized result. +struct ParallelPagedTool { + name: &'static str, + output: String, +} + +impl ToolDefinition for ParallelPagedTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder(self.name) + .description("test tool: returns an oversized parallel result") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .build() + } +} + +#[async_trait] +impl ToolExecutor for ParallelPagedTool { + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + Ok(self.output.clone()) + } +} + +// (a) With paging unconfigured, an oversized result is inserted whole and the +// reader is neither registered nor offered to the model. +#[tokio::test] +async fn unpaged_agents_receive_oversized_results_whole_without_the_reader() { + let model = model_info("model", BuiltinProvider::Anthropic); + let full = numbered_lines("line", 40); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + super::support::tool_use_stream(&model.id, "call-1", "big_tool", r#"{}"#), + super::support::text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(unlimited_results()) + .with_tool(StaticTool::success("big_tool", &full)) + .build() + .expect("build runtime"); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run the big tool")]) + .await + .expect("send"); + + let results = tool_results(agent.history()); + assert_eq!(results.len(), 1); + assert_eq!(results[0].1, full, "the result must be byte-identical"); + assert!( + agent + .tools() + .iter() + .all(|tool| tool.name != "read_tool_result"), + "read_tool_result must not be offered to an unpaged agent" + ); + assert!( + runtime.tool_descriptor("read_tool_result").is_none(), + "read_tool_result must not be registered for an unpaged agent" + ); +} + +// (e) With paging enabled, a result at or below the threshold is still +// byte-identical — only the roster changes. +#[tokio::test] +async fn sub_threshold_results_stay_byte_identical_under_paging() { + let model = model_info("model", BuiltinProvider::Anthropic); + let full = numbered_lines("line", 40); + assert_eq!(full.len(), 800); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + super::support::tool_use_stream(&model.id, "call-1", "big_tool", r#"{}"#), + super::support::text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(unlimited_results()) + .with_tool(StaticTool::success("big_tool", &full)) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("agent", model, paged_config(800, 100)) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run the big tool")]) + .await + .expect("send"); + + let results = tool_results(agent.history()); + assert_eq!(results[0].1, full, "a result at the threshold is not paged"); + assert!(!results[0].1.contains("[paged:")); + assert!( + agent + .tools() + .iter() + .any(|tool| tool.name == "read_tool_result"), + "the reader is offered whenever paging is enabled, not only once it fires" + ); +} + +// (b) An oversized result reaches the model as page 1 plus a trailer, while +// the event stream still carries the complete block. +#[tokio::test] +async fn oversized_results_reach_the_model_paged_and_the_event_stream_whole() { + let model = model_info("model", BuiltinProvider::Anthropic); + let full = numbered_lines("line", 40); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + super::support::tool_use_stream(&model.id, "call-1", "big_tool", r#"{}"#), + super::support::text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(unlimited_results()) + .with_tool(StaticTool::success("big_tool", &full)) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("agent", model, paged_config(100, 100)) + .expect("spawn agent"); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::text("run the big tool")]) + .await + .expect("send"); + + let results = tool_results(agent.history()); + assert_eq!(results.len(), 1); + let page = &results[0].1; + assert!(page.starts_with(&numbered_lines("line", 5))); + assert!(!page.contains("line-006")); + assert!( + page.contains( + "…[paged: lines 1–5 of 40 (0.1 KB of 0.8 KB). \ + Call read_tool_result(tool_use_id=\"call-1\", start_line=6) for the next window.]" + ), + "unexpected trailer: {page}" + ); + + let finished = collect_events(&mut events) + .into_iter() + .find_map(|event| match event { + AgentEvent::ToolExecutionFinished { result } => Some(result), + _ => None, + }) + .expect("a ToolExecutionFinished event"); + let ContentBlock::ToolResult { content, .. } = finished else { + panic!("expected a tool result block"); + }; + assert_eq!( + content.to_display_string(), + full, + "the event stream must keep carrying the unpaged result" + ); +} + +// (c) Successive windows tile the result with absolute line numbers, the last +// one is marked as the end, and a start_line past the end is empty. +#[tokio::test] +async fn read_tool_result_windows_tile_the_result_and_mark_the_end() { + let model = model_info("model", BuiltinProvider::Anthropic); + let full = numbered_lines("line", 40); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + super::support::tool_use_stream(&model.id, "call-1", "big_tool", r#"{}"#), + super::support::tool_use_stream( + &model.id, + "call-2", + "read_tool_result", + &read_window_input("call-1", 6), + ), + super::support::tool_use_stream( + &model.id, + "call-3", + "read_tool_result", + &read_window_input("call-1", 36), + ), + super::support::tool_use_stream( + &model.id, + "call-4", + "read_tool_result", + &read_window_input("call-1", 41), + ), + super::support::text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(unlimited_results()) + .with_tool(StaticTool::success("big_tool", &full)) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("agent", model, paged_config(100, 100)) + .expect("spawn agent"); + + let message = agent + .send(vec![ContentBlock::text("read the whole result")]) + .await + .expect("send"); + assert_eq!(message.text(), "done"); + + let results = tool_results(agent.history()); + assert_eq!(results.len(), 4); + + assert!(results[1].1.starts_with("line-006")); + assert!(results[1].1.contains("lines 6–10 of 40")); + assert!(results[1].1.contains("start_line=11")); + + assert!(results[2].1.starts_with("line-036")); + assert!(results[2].1.contains("line-040")); + assert!( + results[2].1.ends_with("…[end of result]"), + "the window reaching the last line ends the result: {}", + results[2].1 + ); + assert!(!results[2].1.contains("[paged:")); + + assert_eq!( + results[3].1, "…[end of result]", + "a start_line past the end is an empty window, not an error" + ); + assert!(!results[3].2, "reading past the end is not a tool error"); +} + +// The reader's own windows are never paged again, even when `page_bytes` +// exceeds `threshold_bytes` and a window is therefore itself "oversized". +#[tokio::test] +async fn read_tool_result_windows_are_never_paged_recursively() { + let model = model_info("model", BuiltinProvider::Anthropic); + let full = numbered_lines("line", 40); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + super::support::tool_use_stream(&model.id, "call-1", "big_tool", r#"{}"#), + super::support::tool_use_stream( + &model.id, + "call-2", + "read_tool_result", + &read_window_input("call-1", 16), + ), + super::support::text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(unlimited_results()) + .with_tool(StaticTool::success("big_tool", &full)) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("agent", model, paged_config(100, 300)) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("read a large window")]) + .await + .expect("send"); + + let results = tool_results(agent.history()); + let window = &results[1].1; + assert!(window.starts_with("line-016")); + assert!(window.contains("lines 16–30 of 40")); + assert_eq!( + window.matches("[paged:").count(), + 1, + "a window must carry exactly one trailer, never a trailer nested in a page: {window}" + ); + assert!( + !window.contains("call-2"), + "a window must never be re-paged under its own tool_use_id: {window}" + ); +} + +// (d) An unknown tool_use_id is an ordinary tool error and the run continues. +#[tokio::test] +async fn an_unknown_tool_use_id_is_a_tool_error_and_the_run_continues() { + let model = model_info("model", BuiltinProvider::Anthropic); + let full = numbered_lines("line", 40); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + super::support::tool_use_stream(&model.id, "call-1", "big_tool", r#"{}"#), + super::support::tool_use_stream( + &model.id, + "call-2", + "read_tool_result", + &read_window_input("call-does-not-exist", 1), + ), + super::support::text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(unlimited_results()) + .with_tool(StaticTool::success("big_tool", &full)) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("agent", model, paged_config(100, 100)) + .expect("spawn agent"); + + let message = agent + .send(vec![ContentBlock::text( + "read a result that was never paged", + )]) + .await + .expect("an unknown id must not fail the run"); + + assert_eq!(message.text(), "done"); + let results = tool_results(agent.history()); + assert!(results[1].2, "the failed read is an is_error result"); + assert!(results[1].1.contains("no retained result for tool_use_id")); + assert!(results[1].1.contains("call-does-not-exist")); +} + +// (f) A single line longer than a page is the one case that cuts mid-line: +// it cuts on a character boundary and says so. +#[tokio::test] +async fn a_line_longer_than_a_page_hard_cuts_on_a_character_boundary() { + let model = model_info("model", BuiltinProvider::Anthropic); + // 50 four-byte characters: a 200-byte line, plus a short second line. + let full = format!("{}\ntail line\n", "𝄞".repeat(50)); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + super::support::tool_use_stream(&model.id, "call-1", "big_tool", r#"{}"#), + super::support::tool_use_stream( + &model.id, + "call-2", + "read_tool_result", + &read_window_input("call-1", 2), + ), + super::support::text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(unlimited_results()) + .with_tool(StaticTool::success("big_tool", &full)) + .build() + .expect("build runtime"); + let mut agent = runtime + // 102 is not a multiple of the 4-byte character width, so a correct + // cut must round down to 100. + .spawn_with_config("agent", model, paged_config(100, 102)) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run the long-line tool")]) + .await + .expect("send"); + + let results = tool_results(agent.history()); + let page = &results[0].1; + assert!(page.starts_with(&"𝄞".repeat(25))); + assert!(!page.starts_with(&"𝄞".repeat(26))); + assert!( + page.contains("…[line 1 hard-cut at 100 of 201 bytes"), + "unexpected hard-cut marker: {page}" + ); + assert!(page.contains("start_line=2")); + + assert!( + results[1].1.starts_with("tail line\n"), + "the next window resumes at the following whole line: {}", + results[1].1 + ); + assert!(results[1].1.ends_with("…[end of result]")); +} + +// (g) Parallel oversized results page independently, and each one's full text +// is retained under its own tool_use_id. +#[tokio::test] +async fn parallel_oversized_results_page_independently_per_tool_use_id() { + let model = model_info("model", BuiltinProvider::Anthropic); + let first = numbered_lines("aaaa", 40); + let second = numbered_lines("bbbb", 40); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + multi_tool_use_stream( + &model.id, + &[ + ("call-a", "parallel_a", r#"{}"#), + ("call-b", "parallel_b", r#"{}"#), + ], + ), + super::support::tool_use_stream( + &model.id, + "call-c", + "read_tool_result", + &read_window_input("call-b", 6), + ), + super::support::text_stream(&model.id, "done"), + ], + ); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_policy(unlimited_results()) + .with_tool(ParallelPagedTool { + name: "parallel_a", + output: first, + }) + .with_tool(ParallelPagedTool { + name: "parallel_b", + output: second, + }) + .build() + .expect("build runtime"); + let mut agent = runtime + .spawn_with_config("agent", model, paged_config(100, 100)) + .expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("run both parallel tools")]) + .await + .expect("send"); + + let results = tool_results(agent.history()); + assert_eq!(results.len(), 3); + + assert_eq!(results[0].0, "call-a"); + assert!(results[0].1.starts_with("aaaa-001")); + assert!(results[0].1.contains("tool_use_id=\"call-a\"")); + assert!(!results[0].1.contains("bbbb")); + + assert_eq!(results[1].0, "call-b"); + assert!(results[1].1.starts_with("bbbb-001")); + assert!(results[1].1.contains("tool_use_id=\"call-b\"")); + assert!(!results[1].1.contains("aaaa")); + + assert!( + results[2].1.starts_with("bbbb-006"), + "each result is retained under its own id: {}", + results[2].1 + ); + assert!(results[2].1.contains("lines 6–10 of 40")); +} + +fn collect_events(receiver: &mut tokio::sync::broadcast::Receiver) -> Vec { + let mut events = Vec::new(); + while let Ok(event) = receiver.try_recv() { + events.push(event); + } + events +} diff --git a/vendor/mentra/src/agent/wait.rs b/vendor/mentra/src/agent/wait.rs new file mode 100644 index 0000000..af85215 --- /dev/null +++ b/vendor/mentra/src/agent/wait.rs @@ -0,0 +1,136 @@ +use std::{future::Future, path::PathBuf, pin::Pin}; + +use tokio::sync::watch; + +use crate::{error::RuntimeError, runtime::RuntimeHandle, team::TeamMessage}; + +use super::{Agent, AgentSnapshot, AgentStatus}; + +/// Owned future returned by [`Agent`] and [`AgentWaitHandle`] wait helpers. +/// +/// The future does not borrow the agent, so it can be polled concurrently with +/// a call that holds `&mut Agent`, including [`Agent::run`](crate::Agent::run). +pub type AgentWaitFuture = Pin + Send + 'static>>; + +/// Cloneable observation handle for an agent's snapshot and teammate inbox. +#[derive(Clone)] +pub struct AgentWaitHandle { + snapshots: watch::Receiver, + runtime: RuntimeHandle, + team_dir: PathBuf, + agent_name: String, +} + +impl AgentWaitHandle { + /// Resolves with the first current or future snapshot satisfying `predicate`. + /// + /// If every snapshot sender is dropped first, the final published snapshot + /// is returned even when it does not satisfy `predicate`. Dropping the + /// [`Agent`] alone does not close this channel while runtime observers for + /// that agent still own sender clones. + pub fn wait_for_snapshot

(&self, predicate: P) -> AgentWaitFuture + where + P: Fn(&AgentSnapshot) -> bool + Send + 'static, + { + let mut snapshots = self.snapshots.clone(); + Box::pin(async move { + loop { + let snapshot = snapshots.borrow().clone(); + if predicate(&snapshot) { + return snapshot; + } + if snapshots.changed().await.is_err() { + return snapshots.borrow().clone(); + } + } + }) + } + + /// Waits for the relevant run generation to become terminal. + /// + /// If called while a run is active, this waits for that generation. If + /// called while the agent is initially idle or already terminal, it waits + /// for the *next* generation, avoiding an immediate stale return from a + /// previous run. Terminal statuses are `Finished`, `Failed`, and + /// `Interrupted`; the initial `Idle` snapshot is not a completed run. + pub fn wait_until_idle(&self) -> AgentWaitFuture { + let snapshot = self.snapshots.borrow().clone(); + let target_generation = if is_active(&snapshot.status) { + snapshot.run_generation + } else { + snapshot.run_generation.saturating_add(1) + }; + self.wait_for_snapshot(move |snapshot| { + snapshot.run_generation >= target_generation && is_terminal(&snapshot.status) + }) + } + + /// Waits for and consumes the next batch of teammate replies. + /// + /// This is a host-consumption API, not a non-destructive observer. The + /// underlying inbox read moves pending rows to the store's inflight state + /// and resets `pending_team_messages`; the returned messages will therefore + /// not also be injected into a later provider request. Do not race this + /// helper with `Agent::run` reading the same inbox. The next successful run + /// acknowledges inflight rows; a failed run requeues them. + pub fn wait_for_teammate_reply( + &self, + ) -> AgentWaitFuture, RuntimeError>> { + let snapshots = self.clone(); + let runtime = self.runtime.clone(); + let team_dir = self.team_dir.clone(); + let agent_name = self.agent_name.clone(); + Box::pin(async move { + snapshots + .wait_for_snapshot(|snapshot| snapshot.pending_team_messages > 0) + .await; + runtime.read_team_inbox(&team_dir, &agent_name) + }) + } +} + +impl Agent { + /// Returns a cloneable observation handle that does not borrow this agent. + pub fn wait_handle(&self) -> AgentWaitHandle { + AgentWaitHandle { + snapshots: self.watch_snapshot(), + runtime: self.runtime.clone(), + team_dir: self.config.team.team_dir.clone(), + agent_name: self.name.clone(), + } + } + + /// Owned-future convenience for [`AgentWaitHandle::wait_for_snapshot`]. + pub fn wait_for_snapshot

(&self, predicate: P) -> AgentWaitFuture + where + P: Fn(&AgentSnapshot) -> bool + Send + 'static, + { + self.wait_handle().wait_for_snapshot(predicate) + } + + /// Owned-future convenience for [`AgentWaitHandle::wait_until_idle`]. + pub fn wait_until_idle(&self) -> AgentWaitFuture { + self.wait_handle().wait_until_idle() + } + + /// Owned-future convenience for [`AgentWaitHandle::wait_for_teammate_reply`]. + pub fn wait_for_teammate_reply( + &self, + ) -> AgentWaitFuture, RuntimeError>> { + self.wait_handle().wait_for_teammate_reply() + } +} + +fn is_active(status: &AgentStatus) -> bool { + matches!( + status, + AgentStatus::AwaitingModel | AgentStatus::Streaming | AgentStatus::ExecutingTool { .. } + ) +} + +fn is_terminal(status: &AgentStatus) -> bool { + matches!( + status, + AgentStatus::Finished | AgentStatus::Failed(_) | AgentStatus::Interrupted + ) +} diff --git a/vendor/mentra/src/auth.rs b/vendor/mentra/src/auth.rs new file mode 100644 index 0000000..d8c3087 --- /dev/null +++ b/vendor/mentra/src/auth.rs @@ -0,0 +1 @@ +pub mod openai; diff --git a/vendor/mentra/src/auth/openai.rs b/vendor/mentra/src/auth/openai.rs new file mode 100644 index 0000000..b5593cb --- /dev/null +++ b/vendor/mentra/src/auth/openai.rs @@ -0,0 +1,13 @@ +mod client; +mod credential; +mod store; + +pub use client::{ + DEFAULT_AUTH_URL, DEFAULT_CLIENT_ID, DEFAULT_SCOPE, DEFAULT_TOKEN_URL, OpenAIOAuthClient, + OpenAIOAuthError, OpenAITokenSet, PendingAuthorization, +}; +pub use credential::OpenAIOAuthCredentialSource; +pub use store::{ + FileTokenStore, KeychainTokenStore, MemoryTokenStore, PersistentTokenStoreKind, TokenStore, + persistent_token_store, selected_store_kind, +}; diff --git a/vendor/mentra/src/auth/openai/client.rs b/vendor/mentra/src/auth/openai/client.rs new file mode 100644 index 0000000..9592427 --- /dev/null +++ b/vendor/mentra/src/auth/openai/client.rs @@ -0,0 +1,334 @@ +use std::{collections::BTreeMap, net::SocketAddr, time::Duration as StdDuration}; + +use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; +use rand::RngCore; +use reqwest::StatusCode; +use ring::digest::{SHA256, digest}; +use serde::{Deserialize, Serialize}; +use time::{Duration, OffsetDateTime}; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpListener, + time::timeout, +}; +use url::Url; + +pub const DEFAULT_CLIENT_ID: &str = "T19P7LJMcLZgUbhBzA85goHf"; +pub const DEFAULT_AUTH_URL: &str = "https://auth.openai.com/oauth/authorize"; +pub const DEFAULT_TOKEN_URL: &str = "https://auth.openai.com/oauth/token"; +pub const DEFAULT_SCOPE: &str = "openid profile email offline_access"; + +const CALLBACK_PATH: &str = "/callback"; +const CALLBACK_TIMEOUT: StdDuration = StdDuration::from_secs(300); + +#[derive(Debug, thiserror::Error)] +pub enum OpenAIOAuthError { + #[error("http transport error: {0}")] + Transport(#[from] reqwest::Error), + #[error("failed to parse url: {0}")] + Url(#[from] url::ParseError), + #[error("failed to encode or decode json: {0}")] + Json(#[from] serde_json::Error), + #[error("failed to bind loopback callback listener: {0}")] + Io(#[from] std::io::Error), + #[error("oauth endpoint returned HTTP {status}: {body}")] + Http { status: StatusCode, body: String }, + #[error("callback timed out waiting for browser redirect")] + CallbackTimeout, + #[error("callback did not include an authorization code")] + MissingCode, + #[error("callback state mismatch")] + StateMismatch, + #[error("oauth endpoint did not return an API key")] + MissingApiKey, + #[error("no stored OAuth tokens found")] + MissingStoredTokens, + #[error("token store is unsupported on this platform: {0}")] + UnsupportedStore(&'static str), + #[error("credential store command `{command}` failed with status {status}: {stderr}")] + CredentialCommand { + command: &'static str, + status: i32, + stderr: String, + }, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAITokenSet { + pub access_token: String, + pub refresh_token: String, + pub id_token: Option, + pub api_key: Option, + pub expires_at: OffsetDateTime, +} + +impl OpenAITokenSet { + pub fn is_expired(&self, refresh_skew: Duration) -> bool { + self.expires_at <= OffsetDateTime::now_utc() + refresh_skew + } + + pub fn require_api_key(&self) -> Result<&str, OpenAIOAuthError> { + self.api_key + .as_deref() + .ok_or(OpenAIOAuthError::MissingApiKey) + } +} + +#[derive(Debug, Deserialize)] +struct TokenResponse { + access_token: String, + refresh_token: String, + #[serde(default)] + id_token: Option, + #[serde(default)] + api_key: Option, + expires_in_seconds: i64, +} + +impl TokenResponse { + fn into_tokens(self) -> OpenAITokenSet { + OpenAITokenSet { + access_token: self.access_token, + refresh_token: self.refresh_token, + id_token: self.id_token, + api_key: self.api_key, + expires_at: OffsetDateTime::now_utc() + Duration::seconds(self.expires_in_seconds), + } + } +} + +#[derive(Debug, Clone)] +pub struct OpenAIOAuthClient { + client: reqwest::Client, + client_id: String, + auth_url: Url, + token_url: Url, +} + +impl Default for OpenAIOAuthClient { + fn default() -> Self { + Self::new(DEFAULT_CLIENT_ID) + } +} + +impl OpenAIOAuthClient { + pub fn new(client_id: impl Into) -> Self { + Self { + client: reqwest::Client::builder() + .build() + .expect("Failed to build OpenAI OAuth client"), + client_id: client_id.into(), + auth_url: Url::parse(DEFAULT_AUTH_URL).expect("Failed to parse auth url"), + token_url: Url::parse(DEFAULT_TOKEN_URL).expect("Failed to parse token url"), + } + } + + pub async fn start_authorization(&self) -> Result { + let listener = TcpListener::bind("127.0.0.1:0").await?; + let redirect_addr = listener.local_addr()?; + let redirect_uri = loopback_redirect_uri(redirect_addr)?; + let code_verifier = random_base64_url(32); + let code_challenge = pkce_s256(&code_verifier); + let state = random_base64_url(32); + + let mut authorize_url = self.auth_url.clone(); + authorize_url + .query_pairs_mut() + .append_pair("response_type", "code") + .append_pair("code_challenge_method", "S256") + .append_pair("client_id", &self.client_id) + .append_pair("redirect_uri", redirect_uri.as_str()) + .append_pair("code_challenge", &code_challenge) + .append_pair("scope", DEFAULT_SCOPE) + .append_pair("state", &state); + + Ok(PendingAuthorization { + authorize_url, + redirect_uri, + state, + code_verifier, + listener, + }) + } + + pub async fn exchange_code( + &self, + code: &str, + redirect_uri: &Url, + code_verifier: &str, + ) -> Result { + let body = self + .client + .post(self.token_url.clone()) + .form(&BTreeMap::from([ + ("grant_type", "authorization_code"), + ("client_id", self.client_id.as_str()), + ("redirect_uri", redirect_uri.as_str()), + ("code", code), + ("code_verifier", code_verifier), + ])) + .send() + .await?; + + parse_token_response(body).await + } + + pub async fn refresh_tokens( + &self, + refresh_token: &str, + ) -> Result { + let body = self + .client + .post(self.token_url.clone()) + .form(&BTreeMap::from([ + ("grant_type", "refresh_token"), + ("client_id", self.client_id.as_str()), + ("refresh_token", refresh_token), + ])) + .send() + .await?; + + parse_token_response(body).await + } +} + +pub struct PendingAuthorization { + authorize_url: Url, + redirect_uri: Url, + state: String, + code_verifier: String, + listener: TcpListener, +} + +impl PendingAuthorization { + pub fn authorize_url(&self) -> &Url { + &self.authorize_url + } + + pub fn redirect_uri(&self) -> &Url { + &self.redirect_uri + } + + pub async fn complete( + self, + client: &OpenAIOAuthClient, + ) -> Result { + let code = timeout(CALLBACK_TIMEOUT, receive_code(self.listener, &self.state)) + .await + .map_err(|_| OpenAIOAuthError::CallbackTimeout)??; + client + .exchange_code(&code, &self.redirect_uri, &self.code_verifier) + .await + } +} + +async fn parse_token_response( + response: reqwest::Response, +) -> Result { + if !response.status().is_success() { + return Err(OpenAIOAuthError::Http { + status: response.status(), + body: response.text().await.unwrap_or_default(), + }); + } + + let body = response.json::().await?; + Ok(body.into_tokens()) +} + +async fn receive_code( + listener: TcpListener, + expected_state: &str, +) -> Result { + let (mut stream, _) = listener.accept().await?; + let mut buffer = [0_u8; 8192]; + let bytes_read = stream.read(&mut buffer).await?; + let request = String::from_utf8_lossy(&buffer[..bytes_read]); + let path = request + .lines() + .next() + .and_then(|line| line.split_whitespace().nth(1)) + .unwrap_or("/"); + + let callback_url = Url::parse(&format!("http://localhost{path}"))?; + let params: BTreeMap<_, _> = callback_url.query_pairs().into_owned().collect(); + + let response = if params + .get("state") + .is_some_and(|state| state == expected_state) + { + success_response().to_string() + } else { + error_response("State mismatch") + }; + stream.write_all(response.as_bytes()).await?; + stream.shutdown().await?; + + if params + .get("state") + .is_none_or(|state| state != expected_state) + { + return Err(OpenAIOAuthError::StateMismatch); + } + + params + .get("code") + .cloned() + .ok_or(OpenAIOAuthError::MissingCode) +} + +fn loopback_redirect_uri(addr: SocketAddr) -> Result { + Url::parse(&format!( + "http://{}:{}{CALLBACK_PATH}", + addr.ip(), + addr.port() + )) + .map_err(Into::into) +} + +fn pkce_s256(verifier: &str) -> String { + URL_SAFE_NO_PAD.encode(digest(&SHA256, verifier.as_bytes())) +} + +fn random_base64_url(len: usize) -> String { + let mut bytes = vec![0_u8; len]; + rand::rng().fill_bytes(&mut bytes); + URL_SAFE_NO_PAD.encode(bytes) +} + +fn success_response() -> &'static str { + "HTTP/1.1 200 OK\r\nContent-Type: text/html; charset=utf-8\r\n\r\n

Authorization complete

You can return to Mentra.

" +} + +fn error_response(message: &str) -> String { + format!( + "HTTP/1.1 400 Bad Request\r\nContent-Type: text/html; charset=utf-8\r\n\r\n

Authorization failed

{message}

" + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn token_expiry_uses_refresh_skew() { + let tokens = OpenAITokenSet { + access_token: "access".into(), + refresh_token: "refresh".into(), + id_token: None, + api_key: Some("api".into()), + expires_at: OffsetDateTime::now_utc() + Duration::seconds(30), + }; + + assert!(tokens.is_expired(Duration::seconds(60))); + assert!(!tokens.is_expired(Duration::seconds(5))); + } + + #[test] + fn pkce_challenge_is_url_safe() { + let challenge = pkce_s256("test-verifier"); + assert!(!challenge.contains('=')); + assert!(!challenge.contains('+')); + assert!(!challenge.contains('/')); + } +} diff --git a/vendor/mentra/src/auth/openai/credential.rs b/vendor/mentra/src/auth/openai/credential.rs new file mode 100644 index 0000000..30318f4 --- /dev/null +++ b/vendor/mentra/src/auth/openai/credential.rs @@ -0,0 +1,136 @@ +use std::sync::Arc; + +use async_trait::async_trait; +use time::Duration; +use tokio::sync::Mutex; + +use crate::{ + auth::openai::{ + OpenAIOAuthClient, OpenAIOAuthError, OpenAITokenSet, PendingAuthorization, + PersistentTokenStoreKind, TokenStore, persistent_token_store, + }, + provider::openai::OpenAICredentialSource, +}; + +pub struct OpenAIOAuthCredentialSource { + client: OpenAIOAuthClient, + tokens: Mutex, + store: Option>, + refresh_skew: Duration, +} + +impl OpenAIOAuthCredentialSource { + pub fn new(client: OpenAIOAuthClient, tokens: OpenAITokenSet) -> Self { + Self { + client, + tokens: Mutex::new(tokens), + store: None, + refresh_skew: Duration::seconds(60), + } + } + + pub fn with_store(mut self, store: Arc) -> Self { + self.store = Some(store); + self + } + + pub fn with_refresh_skew(mut self, refresh_skew: Duration) -> Self { + self.refresh_skew = refresh_skew; + self + } + + pub fn from_store( + client: OpenAIOAuthClient, + store: Arc, + ) -> Result { + let tokens = store.load()?.ok_or(OpenAIOAuthError::MissingStoredTokens)?; + Ok(Self::new(client, tokens).with_store(store)) + } + + pub fn from_persistent_store( + client: OpenAIOAuthClient, + kind: PersistentTokenStoreKind, + ) -> Result { + Self::from_store(client, persistent_token_store(kind)) + } + + pub fn from_default_persistent_store( + client: OpenAIOAuthClient, + ) -> Result { + Self::from_persistent_store(client, PersistentTokenStoreKind::Auto) + } + + pub async fn from_store_or_authorize( + client: OpenAIOAuthClient, + store: Arc, + on_pending_authorization: F, + ) -> Result + where + F: FnOnce(&PendingAuthorization), + { + match Self::from_store(client.clone(), store.clone()) { + Ok(source) => Ok(source), + Err(OpenAIOAuthError::MissingStoredTokens) => { + let pending = client.start_authorization().await?; + on_pending_authorization(&pending); + let tokens = pending.complete(&client).await?; + store.save(&tokens)?; + Ok(Self::new(client, tokens).with_store(store)) + } + Err(error) => Err(error), + } + } + + pub async fn from_persistent_store_or_authorize( + client: OpenAIOAuthClient, + kind: PersistentTokenStoreKind, + on_pending_authorization: F, + ) -> Result + where + F: FnOnce(&PendingAuthorization), + { + Self::from_store_or_authorize( + client, + persistent_token_store(kind), + on_pending_authorization, + ) + .await + } + + pub async fn from_default_persistent_store_or_authorize( + client: OpenAIOAuthClient, + on_pending_authorization: F, + ) -> Result + where + F: FnOnce(&PendingAuthorization), + { + Self::from_persistent_store_or_authorize( + client, + PersistentTokenStoreKind::Auto, + on_pending_authorization, + ) + .await + } + + async fn current_api_key(&self) -> Result { + let mut tokens = self.tokens.lock().await; + if tokens.is_expired(self.refresh_skew) { + let refreshed = self.client.refresh_tokens(&tokens.refresh_token).await?; + if let Some(store) = &self.store { + store.save(&refreshed)?; + } + *tokens = refreshed; + } + + Ok(tokens.require_api_key()?.to_string()) + } +} + +#[async_trait] +impl OpenAICredentialSource for OpenAIOAuthCredentialSource { + async fn api_key(&self) -> Result { + self.current_api_key() + .await + .map_err(|error| error.to_string()) + } +} diff --git a/vendor/mentra/src/auth/openai/store.rs b/vendor/mentra/src/auth/openai/store.rs new file mode 100644 index 0000000..3347a9b --- /dev/null +++ b/vendor/mentra/src/auth/openai/store.rs @@ -0,0 +1,348 @@ +use std::{ + fs, + io::Write, + path::{Path, PathBuf}, + sync::{Arc, Mutex}, +}; + +#[cfg(target_os = "macos")] +use std::process::Command; + +use directories::BaseDirs; + +use crate::auth::openai::{OpenAIOAuthError, OpenAITokenSet}; + +const DEFAULT_KEYCHAIN_SERVICE: &str = "com.mentra.openai"; +const DEFAULT_KEYCHAIN_ACCOUNT: &str = "default"; +#[cfg(target_os = "macos")] +const KEYCHAIN_NOT_FOUND_EXIT_CODE: i32 = 44; + +pub trait TokenStore: Send + Sync { + fn load(&self) -> Result, OpenAIOAuthError>; + fn save(&self, tokens: &OpenAITokenSet) -> Result<(), OpenAIOAuthError>; + fn clear(&self) -> Result<(), OpenAIOAuthError>; +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum PersistentTokenStoreKind { + Auto, + File, + Keychain, +} + +impl PersistentTokenStoreKind { + pub fn label(self) -> &'static str { + match self { + Self::Auto => "auto", + Self::File => "file", + Self::Keychain => "keychain", + } + } +} + +#[derive(Clone, Default)] +pub struct MemoryTokenStore { + state: Arc>>, +} + +impl MemoryTokenStore { + pub fn new() -> Self { + Self::default() + } +} + +impl TokenStore for MemoryTokenStore { + fn load(&self) -> Result, OpenAIOAuthError> { + Ok(self + .state + .lock() + .expect("memory token store poisoned") + .clone()) + } + + fn save(&self, tokens: &OpenAITokenSet) -> Result<(), OpenAIOAuthError> { + *self.state.lock().expect("memory token store poisoned") = Some(tokens.clone()); + Ok(()) + } + + fn clear(&self) -> Result<(), OpenAIOAuthError> { + *self.state.lock().expect("memory token store poisoned") = None; + Ok(()) + } +} + +#[derive(Debug, Clone)] +pub struct FileTokenStore { + path: PathBuf, +} + +impl FileTokenStore { + pub fn new(path: impl Into) -> Self { + Self { path: path.into() } + } + + pub fn default_path() -> PathBuf { + let base = BaseDirs::new() + .map(|dirs| dirs.data_local_dir().to_path_buf()) + .unwrap_or_else(|| std::env::current_dir().unwrap_or_else(|_| PathBuf::from("."))); + base.join("mentra").join("auth").join("openai.json") + } + + pub fn path(&self) -> &Path { + &self.path + } +} + +impl Default for FileTokenStore { + fn default() -> Self { + Self::new(Self::default_path()) + } +} + +impl TokenStore for FileTokenStore { + fn load(&self) -> Result, OpenAIOAuthError> { + match fs::read_to_string(&self.path) { + Ok(contents) => Ok(Some(serde_json::from_str(&contents)?)), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(None), + Err(error) => Err(OpenAIOAuthError::Io(error)), + } + } + + fn save(&self, tokens: &OpenAITokenSet) -> Result<(), OpenAIOAuthError> { + if let Some(parent) = self.path.parent() { + fs::create_dir_all(parent)?; + } + + #[cfg(unix)] + let mut file = { + use std::os::unix::fs::OpenOptionsExt; + + fs::OpenOptions::new() + .create(true) + .truncate(true) + .write(true) + .mode(0o600) + .open(&self.path)? + }; + + #[cfg(not(unix))] + let mut file = fs::OpenOptions::new() + .create(true) + .truncate(true) + .write(true) + .open(&self.path)?; + + let payload = serde_json::to_vec_pretty(tokens)?; + file.write_all(&payload)?; + file.flush()?; + Ok(()) + } + + fn clear(&self) -> Result<(), OpenAIOAuthError> { + match fs::remove_file(&self.path) { + Ok(()) => Ok(()), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()), + Err(error) => Err(OpenAIOAuthError::Io(error)), + } + } +} + +#[derive(Debug, Clone)] +pub struct KeychainTokenStore { + #[cfg(target_os = "macos")] + service: String, + #[cfg(target_os = "macos")] + account: String, +} + +impl KeychainTokenStore { + #[cfg(target_os = "macos")] + pub fn new(service: impl Into, account: impl Into) -> Self { + Self { + service: service.into(), + account: account.into(), + } + } + + #[cfg(not(target_os = "macos"))] + pub fn new(service: impl Into, account: impl Into) -> Self { + let _ = (service.into(), account.into()); + Self {} + } + + pub fn default_service() -> &'static str { + DEFAULT_KEYCHAIN_SERVICE + } + + pub fn default_account() -> &'static str { + DEFAULT_KEYCHAIN_ACCOUNT + } +} + +impl Default for KeychainTokenStore { + fn default() -> Self { + Self::new(Self::default_service(), Self::default_account()) + } +} + +impl TokenStore for KeychainTokenStore { + fn load(&self) -> Result, OpenAIOAuthError> { + #[cfg(target_os = "macos")] + { + let output = Command::new("security") + .args([ + "find-generic-password", + "-a", + &self.account, + "-s", + &self.service, + "-w", + ]) + .output()?; + + if output.status.success() { + let secret = String::from_utf8_lossy(&output.stdout); + return Ok(Some(serde_json::from_str(secret.trim())?)); + } + + if output.status.code() == Some(KEYCHAIN_NOT_FOUND_EXIT_CODE) { + return Ok(None); + } + + Err(command_error(output, "security")) + } + + #[cfg(not(target_os = "macos"))] + { + let _ = self; + Err(OpenAIOAuthError::UnsupportedStore("keychain")) + } + } + + fn save(&self, tokens: &OpenAITokenSet) -> Result<(), OpenAIOAuthError> { + #[cfg(target_os = "macos")] + { + let payload = serde_json::to_string(tokens)?; + let output = Command::new("security") + .args([ + "add-generic-password", + "-U", + "-a", + &self.account, + "-s", + &self.service, + "-w", + &payload, + ]) + .output()?; + + if output.status.success() { + return Ok(()); + } + + Err(command_error(output, "security")) + } + + #[cfg(not(target_os = "macos"))] + { + let _ = tokens; + Err(OpenAIOAuthError::UnsupportedStore("keychain")) + } + } + + fn clear(&self) -> Result<(), OpenAIOAuthError> { + #[cfg(target_os = "macos")] + { + let output = Command::new("security") + .args([ + "delete-generic-password", + "-a", + &self.account, + "-s", + &self.service, + ]) + .output()?; + + if output.status.success() || output.status.code() == Some(KEYCHAIN_NOT_FOUND_EXIT_CODE) + { + return Ok(()); + } + + Err(command_error(output, "security")) + } + + #[cfg(not(target_os = "macos"))] + { + Err(OpenAIOAuthError::UnsupportedStore("keychain")) + } + } +} + +pub fn persistent_token_store(kind: PersistentTokenStoreKind) -> Arc { + match selected_store_kind(kind) { + PersistentTokenStoreKind::File => Arc::new(FileTokenStore::default()), + PersistentTokenStoreKind::Keychain => Arc::new(KeychainTokenStore::default()), + PersistentTokenStoreKind::Auto => unreachable!("auto should resolve to a concrete store"), + } +} + +pub fn selected_store_kind(kind: PersistentTokenStoreKind) -> PersistentTokenStoreKind { + match kind { + PersistentTokenStoreKind::Auto => { + if cfg!(target_os = "macos") { + PersistentTokenStoreKind::Keychain + } else { + PersistentTokenStoreKind::File + } + } + other => other, + } +} + +#[cfg(target_os = "macos")] +fn command_error(output: std::process::Output, command: &'static str) -> OpenAIOAuthError { + OpenAIOAuthError::CredentialCommand { + command, + status: output.status.code().unwrap_or(-1), + stderr: String::from_utf8_lossy(&output.stderr).trim().to_string(), + } +} + +#[cfg(test)] +mod tests { + use time::{Duration, OffsetDateTime}; + + use super::*; + + #[test] + fn memory_store_round_trips_tokens() { + let store = MemoryTokenStore::new(); + let tokens = OpenAITokenSet { + access_token: "access".into(), + refresh_token: "refresh".into(), + id_token: Some("id".into()), + api_key: Some("api".into()), + expires_at: OffsetDateTime::now_utc() + Duration::seconds(60), + }; + + store.save(&tokens).expect("save tokens"); + assert_eq!( + store + .load() + .expect("load tokens") + .expect("missing tokens") + .api_key, + Some("api".into()) + ); + } + + #[test] + fn auto_store_resolves_to_platform_backend() { + let resolved = selected_store_kind(PersistentTokenStoreKind::Auto); + if cfg!(target_os = "macos") { + assert_eq!(resolved, PersistentTokenStoreKind::Keychain); + } else { + assert_eq!(resolved, PersistentTokenStoreKind::File); + } + } +} diff --git a/vendor/mentra/src/background.rs b/vendor/mentra/src/background.rs new file mode 100644 index 0000000..37c4e2e --- /dev/null +++ b/vendor/mentra/src/background.rs @@ -0,0 +1,417 @@ +mod hook; +mod observer; +mod store; + +use std::{ + collections::HashMap, + path::PathBuf, + sync::{ + Arc, Mutex, + atomic::{AtomicU64, Ordering}, + }, +}; + +use serde::{Deserialize, Serialize}; +use strum::Display; + +use crate::agent::AgentEvent; +use crate::runtime::{ + RuntimeStore, + control::{CommandOutput, CommandRequest, RuntimeExecutor}, +}; + +pub(crate) use hook::BackgroundHookSink; +pub(crate) use observer::{BackgroundObserverSink, BackgroundRegistration}; +pub use store::BackgroundStore; + +const OUTPUT_PREVIEW_MAX_CHARS: usize = 500; +const NOTIFICATION_PENDING: i64 = 0; +const NOTIFICATION_ACKED: i64 = 2; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Display)] +#[strum(serialize_all = "snake_case")] +#[serde(rename_all = "snake_case")] +pub enum BackgroundTaskStatus { + Running, + Finished, + Failed, + Interrupted, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct BackgroundTaskSummary { + pub id: String, + pub command: String, + pub cwd: PathBuf, + pub status: BackgroundTaskStatus, + pub output_preview: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct BackgroundNotification { + pub task_id: String, + pub command: String, + pub cwd: PathBuf, + pub status: BackgroundTaskStatus, + pub output_preview: String, +} + +#[derive(Clone)] +pub(crate) struct BackgroundTaskManager { + inner: Arc, +} + +struct BackgroundTaskManagerInner { + store: Arc, + executor: Arc, + hooks: Arc, + next_task_id: AtomicU64, + state: Mutex, +} + +#[derive(Default)] +struct BackgroundTaskManagerState { + agents: HashMap, +} + +#[derive(Default)] +struct AgentBackgroundState { + tasks: Vec, + observer: Option, +} + +#[derive(Clone)] +struct BackgroundObserver { + sink: Arc, +} + +impl BackgroundTaskManager { + pub(crate) fn new( + store: Arc, + executor: Arc, + hooks: Arc, + ) -> Self { + Self { + inner: Arc::new(BackgroundTaskManagerInner { + store, + executor, + hooks, + next_task_id: AtomicU64::default(), + state: Mutex::new(BackgroundTaskManagerState::default()), + }), + } + } + + pub(crate) fn register_agent(&self, registration: BackgroundRegistration) { + let BackgroundRegistration { agent_id, observer } = registration; + let tasks = { + let mut state = self + .inner + .state + .lock() + .expect("background manager poisoned"); + let agent = state.agents.entry(agent_id.clone()).or_default(); + agent.tasks = self + .inner + .store + .load_background_tasks(&agent_id) + .unwrap_or_default(); + agent.observer = Some(BackgroundObserver { + sink: observer.clone(), + }); + agent.tasks.clone() + }; + + observer.publish_snapshot(&tasks); + } + + pub(crate) fn start_task( + &self, + agent_id: &str, + request: CommandRequest, + ) -> Result { + let task_id = format!( + "bg-{}", + self.inner.next_task_id.fetch_add(1, Ordering::Relaxed) + 1 + ); + let summary = BackgroundTaskSummary { + id: task_id.clone(), + command: request.spec.display().to_string(), + cwd: request.cwd.clone(), + status: BackgroundTaskStatus::Running, + output_preview: None, + }; + let _ = self + .inner + .store + .upsert_background_task(agent_id, &summary, NOTIFICATION_ACKED); + + let (observer, tasks) = { + let mut state = self + .inner + .state + .lock() + .expect("background manager poisoned"); + let agent = state.agents.entry(agent_id.to_string()).or_default(); + agent.tasks.push(summary.clone()); + (agent.observer.clone(), agent.tasks.clone()) + }; + self.publish_observer( + observer, + tasks, + AgentEvent::BackgroundTaskStarted { + task: summary.clone(), + }, + ); + let _ = + self.inner + .hooks + .task_started(agent_id, &summary.id, &summary.command, &summary.cwd); + + let manager = self.clone(); + let agent_id = agent_id.to_string(); + let executor = self.inner.executor.clone(); + tokio::spawn(async move { + let completed = execute_task(task_id, request, executor).await; + manager.finish_task(&agent_id, completed); + }); + + Ok(summary) + } + + pub(crate) fn running_task_count(&self, agent_id: &str) -> usize { + let state = self + .inner + .state + .lock() + .expect("background manager poisoned"); + state + .agents + .get(agent_id) + .map(|agent| { + agent + .tasks + .iter() + .filter(|task| task.status == BackgroundTaskStatus::Running) + .count() + }) + .unwrap_or(0) + } + + pub(crate) fn drain_notifications(&self, agent_id: &str) -> Vec { + self.inner + .store + .drain_background_notifications(agent_id) + .unwrap_or_default() + } + + pub(crate) fn has_pending_notifications(&self, agent_id: &str) -> bool { + self.inner + .store + .has_pending_background_notifications(agent_id) + .unwrap_or(false) + } + + pub(crate) fn has_deliverable_notifications(&self, agent_id: &str) -> bool { + self.inner + .store + .has_deliverable_background_notifications(agent_id) + .unwrap_or(false) + } + + pub(crate) fn requeue_notifications( + &self, + agent_id: &str, + notifications: Vec, + ) { + if notifications.is_empty() { + return; + } + let _ = self.inner.store.requeue_background_notifications(agent_id); + } + + pub(crate) fn acknowledge_notifications(&self, agent_id: &str) { + let _ = self.inner.store.ack_background_notifications(agent_id); + } + + pub(crate) fn check_task( + &self, + agent_id: &str, + task_id: Option<&str>, + ) -> Result { + let state = self + .inner + .state + .lock() + .expect("background manager poisoned"); + let Some(agent) = state.agents.get(agent_id) else { + return Ok("No background tasks.".to_string()); + }; + + if let Some(task_id) = task_id { + let task = agent + .tasks + .iter() + .find(|task| task.id == task_id) + .ok_or_else(|| format!("Unknown background task {task_id}"))?; + return Ok(render_task_detail(task)); + } + + if agent.tasks.is_empty() { + return Ok("No background tasks.".to_string()); + } + + Ok(agent + .tasks + .iter() + .map(render_task_summary) + .collect::>() + .join("\n")) + } + + fn finish_task(&self, agent_id: &str, completed: CompletedBackgroundTask) { + let summary = BackgroundTaskSummary { + id: completed.id.clone(), + command: completed.command.clone(), + cwd: completed.cwd.clone(), + status: completed.status.clone(), + output_preview: Some(completed.output_preview.clone()), + }; + let (observer, tasks) = { + let mut state = self + .inner + .state + .lock() + .expect("background manager poisoned"); + let agent = state.agents.entry(agent_id.to_string()).or_default(); + if let Some(existing) = agent.tasks.iter_mut().find(|task| task.id == summary.id) { + *existing = summary.clone(); + } else { + agent.tasks.push(summary.clone()); + } + (agent.observer.clone(), agent.tasks.clone()) + }; + let _ = self + .inner + .store + .upsert_background_task(agent_id, &summary, NOTIFICATION_PENDING); + let status = summary.status.to_string(); + let _ = self + .inner + .hooks + .task_finished(agent_id, &summary.id, &status); + + self.publish_observer( + observer, + tasks, + AgentEvent::BackgroundTaskFinished { task: summary }, + ); + } + fn publish_observer( + &self, + observer: Option, + tasks: Vec, + event: AgentEvent, + ) { + let Some(observer) = observer else { + return; + }; + + observer.sink.publish_snapshot(&tasks); + observer.sink.publish_event(event); + } +} + +struct CompletedBackgroundTask { + id: String, + command: String, + cwd: PathBuf, + status: BackgroundTaskStatus, + output_preview: String, +} + +async fn execute_task( + id: String, + request: CommandRequest, + executor: Arc, +) -> CompletedBackgroundTask { + let command = request.spec.display().to_string(); + let cwd = request.cwd.clone(); + match executor.run(request).await { + Ok(output) => completed_task_from_output(id, command, cwd, output), + Err(error) => CompletedBackgroundTask { + id, + command, + cwd, + status: BackgroundTaskStatus::Failed, + output_preview: truncate_preview(&error), + }, + } +} + +fn completed_task_from_output( + id: String, + command: String, + cwd: PathBuf, + output: CommandOutput, +) -> CompletedBackgroundTask { + let combined = format!("{} {}", output.stdout, output.stderr); + let preview = if combined.trim().is_empty() { + "(no output)".to_string() + } else { + truncate_preview(&combined) + }; + let status = if output.success() { + BackgroundTaskStatus::Finished + } else { + BackgroundTaskStatus::Failed + }; + + CompletedBackgroundTask { + id, + command, + cwd, + status, + output_preview: preview, + } +} + +fn truncate_preview(text: &str) -> String { + let mut compact = String::new(); + for (index, chunk) in text.split_whitespace().enumerate() { + if index > 0 { + compact.push(' '); + } + compact.push_str(chunk); + } + + let mut truncated = compact + .chars() + .take(OUTPUT_PREVIEW_MAX_CHARS) + .collect::(); + if compact.chars().count() > OUTPUT_PREVIEW_MAX_CHARS { + truncated.push_str("..."); + } + truncated +} + +fn render_task_summary(task: &BackgroundTaskSummary) -> String { + format!( + "{}: [{}] cwd={} {}", + task.id, + task.status, + task.cwd.display(), + task.command + ) +} + +fn render_task_detail(task: &BackgroundTaskSummary) -> String { + let output = task.output_preview.as_deref().unwrap_or("(running)"); + format!( + "[{}] cwd={}\n{}\n{}", + task.status, + task.cwd.display(), + task.command, + output + ) +} diff --git a/vendor/mentra/src/background/hook.rs b/vendor/mentra/src/background/hook.rs new file mode 100644 index 0000000..be215f2 --- /dev/null +++ b/vendor/mentra/src/background/hook.rs @@ -0,0 +1,20 @@ +use std::path::Path; + +use crate::error::RuntimeError; + +pub(crate) trait BackgroundHookSink: Send + Sync { + fn task_started( + &self, + agent_id: &str, + task_id: &str, + command: &str, + cwd: &Path, + ) -> Result<(), RuntimeError>; + + fn task_finished( + &self, + agent_id: &str, + task_id: &str, + status: &str, + ) -> Result<(), RuntimeError>; +} diff --git a/vendor/mentra/src/background/observer.rs b/vendor/mentra/src/background/observer.rs new file mode 100644 index 0000000..540d1d5 --- /dev/null +++ b/vendor/mentra/src/background/observer.rs @@ -0,0 +1,16 @@ +use std::sync::Arc; + +use crate::agent::AgentEvent; + +use super::BackgroundTaskSummary; + +pub(crate) trait BackgroundObserverSink: Send + Sync { + fn publish_snapshot(&self, tasks: &[BackgroundTaskSummary]); + fn publish_event(&self, event: AgentEvent); +} + +#[derive(Clone)] +pub(crate) struct BackgroundRegistration { + pub(crate) agent_id: String, + pub(crate) observer: Arc, +} diff --git a/vendor/mentra/src/background/store.rs b/vendor/mentra/src/background/store.rs new file mode 100644 index 0000000..2b40539 --- /dev/null +++ b/vendor/mentra/src/background/store.rs @@ -0,0 +1,27 @@ +use crate::error::RuntimeError; + +use super::{BackgroundNotification, BackgroundTaskSummary}; + +pub trait BackgroundStore: Send + Sync { + fn load_background_tasks( + &self, + agent_id: &str, + ) -> Result, RuntimeError>; + fn upsert_background_task( + &self, + agent_id: &str, + task: &BackgroundTaskSummary, + notification_state: i64, + ) -> Result<(), RuntimeError>; + fn drain_background_notifications( + &self, + agent_id: &str, + ) -> Result, RuntimeError>; + fn has_deliverable_background_notifications( + &self, + agent_id: &str, + ) -> Result; + fn has_pending_background_notifications(&self, agent_id: &str) -> Result; + fn ack_background_notifications(&self, agent_id: &str) -> Result<(), RuntimeError>; + fn requeue_background_notifications(&self, agent_id: &str) -> Result<(), RuntimeError>; +} diff --git a/vendor/mentra/src/compaction.rs b/vendor/mentra/src/compaction.rs new file mode 100644 index 0000000..7bec6e5 --- /dev/null +++ b/vendor/mentra/src/compaction.rs @@ -0,0 +1,775 @@ +#[cfg(test)] +mod tests; + +use std::{ + borrow::Cow, + collections::HashSet, + path::Path, + path::PathBuf, + sync::Arc, + time::{SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use regex::Regex; + +use crate::{ + ContentBlock, Message, + error::RuntimeError, + provider::{ + CompactionInputItem, CompactionRequest as ProviderCompactionRequest, + CompactionResponse as ProviderCompactionResponse, Provider, ProviderError, + ProviderRequestOptions, Request, + }, + transcript::{AgentTranscript, CompactionSummary, TranscriptItem, TranscriptKind}, +}; + +/// Context mechanically extracted from transcript items before summarization. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct ExtractedContext { + pub files_touched: Vec, + pub verification_outcomes: Vec, + pub permission_decisions: Vec, +} + +/// Scan transcript items to extract file paths, verification outcomes, and permission decisions. +pub fn extract_context(items: &[TranscriptItem]) -> ExtractedContext { + use std::sync::LazyLock; + + static FILE_RE: LazyLock = LazyLock::new(|| { + Regex::new(r#"(?:^|[\s"'`(,])([a-zA-Z0-9_.][a-zA-Z0-9_./\-]*\.[a-zA-Z]{1,10})"#) + .expect("valid regex literal") + }); + static VERIFICATION_RE: LazyLock = LazyLock::new(|| { + Regex::new( + r"(?i)(cargo\s+test|pytest|npm\s+test|jest|mocha|go\s+test|make\s+test|rspec|yarn\s+test).*?(pass|fail|error|ok|success|FAILED|PASSED)", + ) + .expect("valid regex literal") + }); + static PERMISSION_RE: LazyLock = LazyLock::new(|| { + Regex::new(r"(?i)(permission|allowed|denied|approved|rejected|authorized)") + .expect("valid regex literal") + }); + let file_re = &*FILE_RE; + let verification_re = &*VERIFICATION_RE; + let permission_re = &*PERMISSION_RE; + + let mut files_seen = HashSet::new(); + let mut files = Vec::new(); + let mut verifications = Vec::new(); + let mut permissions = Vec::new(); + + for item in items { + let text = item.text(); + let is_tool_exchange = matches!(item.kind, TranscriptKind::ToolExchange { .. }); + + if is_tool_exchange { + for cap in file_re.captures_iter(&text) { + if let Some(m) = cap.get(1) { + let path = m.as_str().to_string(); + if files_seen.insert(path.clone()) { + files.push(path); + } + } + } + } + + for line in text.lines() { + if verification_re.is_match(line) { + let trimmed = line.trim().to_string(); + if !trimmed.is_empty() { + verifications.push(trimmed); + } + } + if permission_re.is_match(line) { + let trimmed = line.trim().to_string(); + if !trimmed.is_empty() { + permissions.push(trimmed); + } + } + } + } + + ExtractedContext { + files_touched: files, + verification_outcomes: verifications, + permission_decisions: permissions, + } +} + +/// Format extracted context as a text preamble for the compaction prompt. +pub fn format_extracted_context(ctx: &ExtractedContext) -> String { + let mut sections = Vec::new(); + + if !ctx.files_touched.is_empty() { + let mut section = String::from("FILES TOUCHED (must preserve):\n"); + for f in &ctx.files_touched { + section.push_str("- "); + section.push_str(f); + section.push('\n'); + } + sections.push(section); + } + + if !ctx.verification_outcomes.is_empty() { + let mut section = String::from("VERIFICATION OUTCOMES (must preserve):\n"); + for v in &ctx.verification_outcomes { + section.push_str("- "); + section.push_str(v); + section.push('\n'); + } + sections.push(section); + } + + if !ctx.permission_decisions.is_empty() { + let mut section = String::from("PERMISSION DECISIONS (must preserve):\n"); + for p in &ctx.permission_decisions { + section.push_str("- "); + section.push_str(p); + section.push('\n'); + } + sections.push(section); + } + + sections.join("\n") +} + +/// Diagnostics captured during a compaction operation. +#[derive(Debug, Clone)] +pub struct CompactionDiagnostics { + pub items_before: usize, + pub items_after: usize, + pub approx_tokens_before: usize, + pub approx_tokens_after: usize, + pub preserved_user_turns: usize, + pub preserved_delegation_results: usize, + pub extracted_facts_count: usize, + pub summary_preview: String, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum CompactionMode { + #[default] + LocalOnly, + PreferRemote, + RemoteOnly, +} + +#[derive(Debug, Clone)] +pub struct CompactionRequest { + pub model: String, + pub transcript: AgentTranscript, + pub transcript_dir: PathBuf, + pub summary_max_input_chars: usize, + pub summary_max_output_tokens: u32, + pub preserve_recent_user_tokens: usize, + pub preserve_recent_delegation_results: usize, + pub provider_request_options: ProviderRequestOptions, + pub mode: CompactionMode, + pub max_persisted_transcripts: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum CompactionExecutionMode { + Local, + Remote, +} + +/// The result of one compaction: a replacement [`AgentTranscript`] plus +/// counts describing how it was built from the original. +/// +/// **Metadata-preservation guarantee (mentra ADR-0001 §6).** Every item that +/// survives into `transcript` — the untouched continuation tail, salvaged +/// recent user turns, and salvaged recent delegation results — is copied +/// **verbatim** from the original transcript, so any opaque +/// [`TranscriptItem::details`] a host attached survives bit-for-bit. This +/// holds by construction: mentra never rebuilds a preserved item from its +/// projected [`crate::Message`] (which would drop `details`, a field that +/// exists only on `TranscriptItem`); it clones the original item. +/// +/// This guarantee is scoped to *preserved* items only. An item inside the +/// summarized prefix that is **not** salvaged is replaced by the +/// [`CompactionSummary`] along with the rest of its content — its `details` +/// go with it. That is honest, documented behavior, not a violation: the +/// contract never promises to resurrect a discarded item's metadata, only to +/// never silently drop it from one that was kept. The full pre-compaction +/// transcript — discarded items included — is written to `transcript_path` +/// before summarization runs, so a host that needs a discarded item's +/// `details` after the fact can still recover them from that snapshot. +#[derive(Debug, Clone)] +pub struct CompactionOutcome { + pub mode: CompactionExecutionMode, + /// Path to the `.jsonl` snapshot of the **entire pre-compaction** + /// transcript, one [`TranscriptItem`] per line, written before + /// summarization runs. Every item's `details` round-trips through this + /// file, including items the compacted `transcript` goes on to discard — + /// this is the recovery artifact for the summarized prefix. + pub transcript_path: PathBuf, + pub transcript: AgentTranscript, + pub summary: CompactionSummary, + /// Count of original items in the summarized prefix + /// (`required_tail_start_for_continuation`'s `preserve_from` split). + /// This counts every item in that prefix, **including** ones also + /// salvaged into `transcript` by `preserved_user_turns` / + /// `preserved_delegation_results` — from this count's point of view they + /// were replaced by the summary, even though their content (and + /// `details`) survives verbatim elsewhere in the replacement transcript. + /// It does not mean "gone". + pub replaced_items: usize, + /// Count of items kept strictly because they are the untouched + /// continuation tail (outside the summarized prefix), independent of + /// `preserved_user_turns` / `preserved_delegation_results` below, which + /// count salvaged items pulled *out of* the summarized prefix instead. + /// The three counts are disjoint by construction. + pub preserved_items: usize, + /// Count of recent user turns salvaged out of the summarized prefix and + /// copied verbatim (details included) into the replacement transcript. + pub preserved_user_turns: usize, + /// Count of recent delegation results salvaged out of the summarized + /// prefix and copied verbatim (details included) into the replacement + /// transcript. + pub preserved_delegation_results: usize, + pub diagnostics: CompactionDiagnostics, +} + +/// Compacts an agent transcript into a shorter one carrying a summary of the +/// discarded portion. See [`CompactionOutcome`] for the metadata-preservation +/// contract every implementation must uphold: `details` on any item that +/// survives compaction (tail, salvaged user turns, salvaged delegation +/// results) is preserved bit-for-bit; `details` on a discarded, unsalvaged +/// item is honestly gone with the rest of that item's content, recoverable +/// only from the pre-compaction snapshot at [`CompactionOutcome::transcript_path`]. +#[async_trait] +pub trait CompactionEngine: Send + Sync { + async fn compact( + &self, + provider: Arc, + request: CompactionRequest, + ) -> Result, RuntimeError>; +} + +/// The default [`CompactionEngine`]: summarizes the compactable prefix of a +/// transcript (locally via the provider's chat completion, or remotely via +/// [`Provider::compact`] when supported), while keeping the continuation +/// tail and a bounded number of recent user turns and delegation results +/// verbatim — including their opaque `details` — per the +/// [`CompactionOutcome`] contract. +#[derive(Debug, Default)] +pub struct StandardCompactionEngine; + +#[async_trait] +impl CompactionEngine for StandardCompactionEngine { + async fn compact( + &self, + provider: Arc, + request: CompactionRequest, + ) -> Result, RuntimeError> { + if request.transcript.is_empty() { + return Ok(None); + } + + let items = request.transcript.items(); + let protected_tail_start = required_tail_start_for_continuation(items); + + // When the protected tail *is* the whole transcript there is nothing + // older to summarize, and compaction used to give up — leaving an + // over-budget turn with no way out. It can be summarized as a unit, + // but only when the thing pinning the tail is a tool call and its + // result: those must travel together, and summarizing both at once + // orphans nothing. + // + // A bare user turn is deliberately excluded. Replacing the user's + // actual instruction with a summary of itself, before the model has + // even read it, loses the very thing the turn exists to convey. + let split_turn = protected_tail_start == 0 && ends_in_tool_exchange(items); + let (compacted_prefix, tail_start) = if split_turn { + (items, items.len()) + } else { + (&items[..protected_tail_start], protected_tail_start) + }; + if compacted_prefix.is_empty() { + return Ok(None); + } + + let transcript_path = + persist_transcript(request.transcript.items(), &request.transcript_dir).await?; + if let Some(max) = request.max_persisted_transcripts { + let _ = cleanup_old_transcripts(&request.transcript_dir, max).await; + } + let supports_remote = provider.capabilities().supports_history_compaction; + let (mode, mut summary) = match request.mode { + CompactionMode::LocalOnly => ( + CompactionExecutionMode::Local, + summarize_locally(provider, &request, compacted_prefix).await?, + ), + CompactionMode::PreferRemote => { + if supports_remote { + match compact_remotely(provider.clone(), &request, compacted_prefix).await { + Ok(Some(summary)) => (CompactionExecutionMode::Remote, summary), + Ok(None) + | Err(RuntimeError::FailedToCompactHistory( + ProviderError::UnsupportedCapability(_), + )) => ( + CompactionExecutionMode::Local, + summarize_locally(provider, &request, compacted_prefix).await?, + ), + Err(error) => return Err(error), + } + } else { + ( + CompactionExecutionMode::Local, + summarize_locally(provider, &request, compacted_prefix).await?, + ) + } + } + CompactionMode::RemoteOnly => { + if !supports_remote { + return Err(RuntimeError::FailedToCompactHistory( + ProviderError::UnsupportedCapability("history_compaction".to_string()), + )); + } + ( + CompactionExecutionMode::Remote, + compact_remotely(provider, &request, compacted_prefix) + .await? + .ok_or_else(|| { + RuntimeError::FailedToCompactHistory( + ProviderError::UnsupportedCapability( + "history_compaction".to_string(), + ), + ) + })?, + ) + } + }; + + let items_before = request.transcript.len(); + let tokens_before = approx_token_count_items(request.transcript.items()); + + let preserved_user_turns = + select_recent_user_turns(compacted_prefix, request.preserve_recent_user_tokens); + let preserved_delegation_results = select_recent_delegation_results( + compacted_prefix, + request.preserve_recent_delegation_results, + ); + + let extracted = extract_context(compacted_prefix); + let extracted_facts_count = extracted.files_touched.len() + + extracted.verification_outcomes.len() + + extracted.permission_decisions.len(); + + // Union with whatever the previous compaction recorded, so the set + // grows monotonically instead of being re-derived from a prefix that + // no longer contains the older tool exchanges. + summary.files_touched = accumulate_files( + carried_files(items), + extracted.files_touched.iter().map(String::as_str), + ); + + let mut replacement = Vec::new(); + replacement.extend(preserved_user_turns.iter().cloned()); + for item in &preserved_delegation_results { + if !replacement.contains(item) { + replacement.push(item.clone()); + } + } + replacement.push(TranscriptItem::compaction_summary(summary.clone())); + replacement.extend_from_slice(&items[tail_start..]); + + let items_after = replacement.len(); + let tokens_after = approx_token_count_items(&replacement); + + let summary_preview = summary + .render_for_handoff() + .chars() + .take(200) + .collect::(); + + let diagnostics = CompactionDiagnostics { + items_before, + items_after, + approx_tokens_before: tokens_before, + approx_tokens_after: tokens_after, + preserved_user_turns: preserved_user_turns.len(), + preserved_delegation_results: preserved_delegation_results.len(), + extracted_facts_count, + summary_preview, + }; + + Ok(Some(CompactionOutcome { + mode, + transcript_path, + transcript: AgentTranscript::new(replacement), + summary, + replaced_items: compacted_prefix.len(), + preserved_items: request.transcript.len().saturating_sub(tail_start), + preserved_user_turns: preserved_user_turns.len(), + preserved_delegation_results: preserved_delegation_results.len(), + diagnostics, + })) + } +} + +pub(crate) fn compaction_request_from_agent( + model: &str, + transcript: AgentTranscript, + config: &crate::agent::CompactionConfig, + provider_request_options: ProviderRequestOptions, +) -> CompactionRequest { + CompactionRequest { + model: model.to_string(), + transcript, + transcript_dir: config.transcript_dir.clone(), + summary_max_input_chars: config.summary_max_input_chars, + summary_max_output_tokens: config.summary_max_output_tokens, + preserve_recent_user_tokens: config.preserve_recent_user_tokens, + preserve_recent_delegation_results: config.preserve_recent_delegation_results, + provider_request_options, + mode: config.mode, + max_persisted_transcripts: config.max_persisted_transcripts, + } +} + +async fn summarize_locally( + provider: Arc, + request: &CompactionRequest, + items: &[TranscriptItem], +) -> Result { + let summary_items = items_without_thinking(items); + let serialized = + serde_json::to_string(&summary_items).map_err(RuntimeError::FailedToSerializeTranscript)?; + let transcript = truncate_to_char_boundary(&serialized, request.summary_max_input_chars); + + let extracted = extract_context(items); + let context_preamble = format_extracted_context(&extracted); + + let system = "\ +You are a coding-session compaction engine. Your job is to compress an agent transcript \ +into a structured JSON summary that preserves all operationally critical context for \ +session continuity.\n\n\ +You MUST preserve:\n\ +- All file paths that were read, written, or modified\n\ +- Shell command outcomes (build results, test pass/fail, lint output)\n\ +- Permission decisions (what was allowed, denied, or deferred)\n\ +- Architectural decisions and their rationale\n\ +- Constraints and invariants discovered during the session\n\ +- Current working state (what is done, what is in progress, what remains)\n\ +- Error states and how they were resolved\n\ +- Delegated work outcomes and pending delegations\n\n\ +Return strict JSON with keys: goal, progress, decisions, constraints, \ +delegated_work, artifacts, open_questions, next_steps.\n\ +Each key should contain concrete, specific information -- not vague summaries.\n\ +File paths, command outputs, and error messages should be quoted verbatim."; + + let mut prompt = String::new(); + if !context_preamble.is_empty() { + prompt.push_str("=== EXTRACTED FACTS (must preserve verbatim) ===\n"); + prompt.push_str(&context_preamble); + prompt.push_str("\n=== END EXTRACTED FACTS ===\n\n"); + } + prompt.push_str("Summarize this agent transcript for continuity and multi-agent handoff. Preserve goal, progress, concrete decisions, constraints, delegated work outcomes, artifacts, open questions, and next steps.\n\nTranscript JSON:\n"); + prompt.push_str(transcript); + let response = provider + .send(Request { + model: Cow::Borrowed(request.model.as_str()), + system: Some(Cow::Borrowed(system)), + messages: Cow::Owned(vec![Message::user(ContentBlock::text(prompt))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: Some(request.summary_max_output_tokens), + metadata: Cow::Owned(Default::default()), + provider_request_options: request.provider_request_options.clone(), + }) + .await + .map_err(RuntimeError::FailedToCompactHistory)?; + let text = response + .content + .into_iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text), + _ => None, + }) + .collect::>() + .join("\n") + .trim() + .to_string(); + if text.is_empty() { + return Ok(CompactionSummary::default()); + } + + serde_json::from_str(&text) + .unwrap_or_else(|_| CompactionSummary::from_fallback_text(text)) + .pipe(Ok) +} + +fn items_without_thinking(items: &[TranscriptItem]) -> Vec { + items + .iter() + .cloned() + .map(|mut item| { + if let Some(message) = item.message.as_mut() { + message + .content + .retain(|block| !matches!(block, ContentBlock::Thinking { .. })); + } + item + }) + .collect() +} + +async fn compact_remotely( + provider: Arc, + request: &CompactionRequest, + items: &[TranscriptItem], +) -> Result, RuntimeError> { + let input = items + .iter() + .map(project_compaction_item) + .collect::>(); + let response = provider + .compact(ProviderCompactionRequest { + model: Cow::Borrowed(request.model.as_str()), + instructions: Cow::Borrowed( + "Compact this transcript into a continuity handoff that preserves delegated work.", + ), + input: Cow::Owned(input), + metadata: Cow::Owned(Default::default()), + provider_request_options: request.provider_request_options.clone(), + }) + .await + .map_err(RuntimeError::FailedToCompactHistory)?; + Ok(parse_remote_summary(response)) +} + +fn parse_remote_summary(response: ProviderCompactionResponse) -> Option { + response + .output + .into_iter() + .rev() + .find_map(|item| match item { + CompactionInputItem::CompactionSummary { content } => serde_json::from_str(&content) + .ok() + .or_else(|| Some(CompactionSummary::from_fallback_text(content))), + _ => None, + }) +} + +fn project_compaction_item(item: &TranscriptItem) -> CompactionInputItem { + match &item.kind { + TranscriptKind::UserTurn => CompactionInputItem::UserTurn { + content: item.text(), + }, + TranscriptKind::AssistantTurn => CompactionInputItem::AssistantTurn { + content: item.text(), + }, + TranscriptKind::ToolExchange { is_error, .. } => CompactionInputItem::ToolExchange { + request: None, + result: item.text(), + is_error: *is_error, + }, + TranscriptKind::CanonicalContext => CompactionInputItem::CanonicalContext { + content: item.text(), + }, + TranscriptKind::MemoryRecall => CompactionInputItem::MemoryRecall { + content: item.text(), + }, + TranscriptKind::DelegationRequest { delegation, .. } + | TranscriptKind::DelegationResult { delegation, .. } => { + CompactionInputItem::DelegationResult { + agent_id: delegation.agent_id.clone(), + agent_name: delegation.agent_name.clone(), + role: delegation.role.clone(), + status: format!("{:?}", delegation.status).to_lowercase(), + content: item.text(), + } + } + TranscriptKind::CompactionSummary { summary } => CompactionInputItem::CompactionSummary { + content: summary.render_for_handoff(), + }, + } +} + +/// Whether the transcript ends with a tool exchange, i.e. the pinned tail is +/// a tool call and its result rather than a plain turn. +fn ends_in_tool_exchange(items: &[TranscriptItem]) -> bool { + items + .last() + .is_some_and(|item| matches!(item.kind, TranscriptKind::ToolExchange { .. })) +} + +/// The file list recorded by the newest compaction already in `items`. +/// +/// `extract_context` only scans tool exchanges, and a compaction summary is +/// not one, so without this the previous round's findings would be invisible +/// to the next. +fn carried_files(items: &[TranscriptItem]) -> Vec { + items + .iter() + .rev() + .find_map(|item| match &item.kind { + TranscriptKind::CompactionSummary { summary } => Some(summary.files_touched.clone()), + _ => None, + }) + .unwrap_or_default() +} + +/// Unions two file lists, preserving first-seen order and dropping repeats. +fn accumulate_files<'a>(carried: Vec, fresh: impl Iterator) -> Vec { + let mut seen: HashSet = carried.iter().cloned().collect(); + let mut files = carried; + for path in fresh { + if seen.insert(path.to_string()) { + files.push(path.to_string()); + } + } + files +} + +fn select_recent_user_turns(items: &[TranscriptItem], token_budget: usize) -> Vec { + let mut selected = Vec::new(); + let mut remaining = token_budget; + for item in items.iter().rev() { + if !item.is_real_user_turn() { + continue; + } + let tokens = approx_token_count(&item.text()); + if tokens > remaining && !selected.is_empty() { + break; + } + remaining = remaining.saturating_sub(tokens); + selected.push(item.clone()); + if remaining == 0 { + break; + } + } + selected.reverse(); + selected +} + +fn select_recent_delegation_results( + items: &[TranscriptItem], + max_items: usize, +) -> Vec { + let mut selected = items + .iter() + .filter(|item| item.is_delegation_result()) + .rev() + .take(max_items) + .cloned() + .collect::>(); + selected.reverse(); + selected +} + +fn required_tail_start_for_continuation(items: &[TranscriptItem]) -> usize { + let Some(last_index) = items.len().checked_sub(1) else { + return 0; + }; + let last = &items[last_index]; + if matches!(last.kind, TranscriptKind::ToolExchange { .. }) + && last_index > 0 + && matches!(items[last_index - 1].kind, TranscriptKind::AssistantTurn) + { + last_index - 1 + } else { + last_index + } +} + +fn approx_token_count(text: &str) -> usize { + let char_estimate = text.chars().count().div_ceil(4); + let word_count = text.split_whitespace().count(); + let word_estimate = ((word_count as f64) * 1.3).ceil() as usize; + char_estimate.max(word_estimate) +} + +fn approx_token_count_items(items: &[TranscriptItem]) -> usize { + items + .iter() + .map(|item| approx_token_count(&item.text())) + .sum() +} + +async fn persist_transcript( + transcript: &[TranscriptItem], + transcript_dir: &Path, +) -> Result { + tokio::fs::create_dir_all(transcript_dir) + .await + .map_err(RuntimeError::FailedToPersistTranscript)?; + + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time should be after unix epoch") + .as_nanos(); + let transcript_path = transcript_dir.join(format!("{timestamp}.jsonl")); + let mut serialized = String::new(); + for item in transcript { + let line = + serde_json::to_string(item).map_err(RuntimeError::FailedToSerializeTranscript)?; + serialized.push_str(&line); + serialized.push('\n'); + } + tokio::fs::write(&transcript_path, serialized) + .await + .map_err(RuntimeError::FailedToPersistTranscript)?; + Ok(transcript_path) +} + +/// Removes the oldest transcript files in `dir` when count exceeds `keep`. +/// Files are sorted by filename (nanosecond timestamps → oldest first). +/// Delete errors are ignored — this is best-effort cleanup. +pub(crate) async fn cleanup_old_transcripts(dir: &Path, keep: usize) -> Result<(), RuntimeError> { + let mut read_dir = tokio::fs::read_dir(dir) + .await + .map_err(RuntimeError::FailedToPersistTranscript)?; + + let mut files: Vec = Vec::new(); + while let Some(entry) = read_dir + .next_entry() + .await + .map_err(RuntimeError::FailedToPersistTranscript)? + { + let path = entry.path(); + if path.extension().and_then(|e| e.to_str()) == Some("jsonl") { + files.push(path); + } + } + + if files.len() <= keep { + return Ok(()); + } + + // Sort ascending by filename — nanosecond timestamps put oldest first. + files.sort_by(|a, b| a.file_name().cmp(&b.file_name())); + + let to_delete = files.len() - keep; + for path in files.iter().take(to_delete) { + let _ = tokio::fs::remove_file(path).await; + } + + Ok(()) +} + +fn truncate_to_char_boundary(input: &str, max_chars: usize) -> &str { + if input.chars().count() <= max_chars { + return input; + } + + let mut end = input.len(); + for (index, _) in input.char_indices().take(max_chars + 1) { + end = index; + } + &input[..end] +} + +trait Pipe: Sized { + fn pipe(self, f: impl FnOnce(Self) -> T) -> T { + f(self) + } +} + +impl Pipe for T {} diff --git a/vendor/mentra/src/compaction/tests.rs b/vendor/mentra/src/compaction/tests.rs new file mode 100644 index 0000000..026ca48 --- /dev/null +++ b/vendor/mentra/src/compaction/tests.rs @@ -0,0 +1,678 @@ +use std::{ + collections::BTreeMap, + sync::atomic::{AtomicU64, Ordering}, +}; + +use serde_json::{Value, json}; + +use super::*; +use crate::{ + ContentBlock, DelegationArtifact, DelegationKind, DelegationStatus, Message, ModelInfo, Role, + provider::{ + ProviderDescriptor, ProviderEventStream, Response, provider_event_stream_from_response, + }, +}; + +/// Asserts an entry survived compaction unchanged in identity and content. +/// +/// `parent_id` is deliberately excluded: it records where an entry sits on the +/// active path, and compaction genuinely moves salvaged entries. Holding the +/// old link would leave it pointing at an entry the replacement transcript no +/// longer contains. Identity (`id`), content, and `details` must not move. +fn assert_same_entry(actual: &TranscriptItem, expected: &TranscriptItem, context: &str) { + assert_eq!(actual.id, expected.id, "{context} (identity)"); + assert_eq!(actual.kind, expected.kind, "{context} (kind)"); + assert_eq!(actual.message, expected.message, "{context} (message)"); + assert_eq!(actual.details(), expected.details(), "{context} (details)"); +} + +fn tool_exchange_item(text: &str) -> TranscriptItem { + TranscriptItem::tool_exchange( + Message::user(ContentBlock::text(text)), + Some("tool_1".to_string()), + false, + ) +} + +fn user_turn_item(text: &str) -> TranscriptItem { + TranscriptItem::user_turn(Message::user(ContentBlock::text(text))) +} + +#[test] +fn local_summary_projection_excludes_thinking_from_full_transcript_json() { + let items = vec![TranscriptItem::assistant_turn(Message { + role: Role::Assistant, + content: vec![ + ContentBlock::Thinking { + thinking: "private chain".to_string(), + signature: Some("opaque-signature".to_string()), + encrypted_content: None, + id: None, + provenance: Some(crate::ReasoningProvenance { + provider: crate::ProviderId::new("anthropic"), + model: "claude-test".to_string(), + format: crate::ReasoningFormat::AnthropicSigned, + }), + redacted: false, + }, + ContentBlock::text("visible answer"), + ], + })]; + + let serialized = serde_json::to_string(&items_without_thinking(&items)).unwrap(); + + assert!(serialized.contains("visible answer")); + assert!(!serialized.contains("private chain")); + assert!(!serialized.contains("opaque-signature")); + assert!(!serialized.contains("Thinking")); + assert_eq!(items[0].message.as_ref().unwrap().content.len(), 2); +} + +#[test] +fn extract_context_finds_file_paths_in_tool_exchanges() { + let items = vec![ + tool_exchange_item("Reading file src/main.rs and also lib/utils.py"), + tool_exchange_item("Modified path/to/config.toml successfully"), + ]; + let ctx = extract_context(&items); + assert!( + ctx.files_touched.contains(&"src/main.rs".to_string()), + "should find src/main.rs, got: {:?}", + ctx.files_touched + ); + assert!( + ctx.files_touched.contains(&"lib/utils.py".to_string()), + "should find lib/utils.py, got: {:?}", + ctx.files_touched + ); + assert!( + ctx.files_touched + .contains(&"path/to/config.toml".to_string()), + "should find path/to/config.toml, got: {:?}", + ctx.files_touched + ); +} + +#[test] +fn extract_context_deduplicates_file_paths() { + let items = vec![ + tool_exchange_item("Reading src/main.rs"), + tool_exchange_item("Writing src/main.rs again"), + ]; + let ctx = extract_context(&items); + let count = ctx + .files_touched + .iter() + .filter(|p| p.as_str() == "src/main.rs") + .count(); + assert_eq!(count, 1, "file paths should be deduplicated"); +} + +#[test] +fn extract_context_ignores_file_paths_in_non_tool_items() { + let items = vec![user_turn_item("Please edit src/main.rs")]; + let ctx = extract_context(&items); + assert!( + ctx.files_touched.is_empty(), + "user turns should not contribute file paths, got: {:?}", + ctx.files_touched + ); +} + +#[test] +fn extract_context_finds_verification_outcomes() { + let items = vec![ + tool_exchange_item("Running: cargo test result: ok. 5 passed; 0 FAILED"), + tool_exchange_item("npm test completed with error code 1"), + ]; + let ctx = extract_context(&items); + assert!( + !ctx.verification_outcomes.is_empty(), + "should find verification outcomes" + ); + assert!( + ctx.verification_outcomes + .iter() + .any(|v| v.contains("cargo test") || v.contains("FAILED")), + "should find cargo test outcome, got: {:?}", + ctx.verification_outcomes + ); +} + +#[test] +fn extract_context_finds_verification_in_any_item_kind() { + let items = vec![user_turn_item("cargo test result: 10 passed; 0 FAILED")]; + let ctx = extract_context(&items); + assert!( + !ctx.verification_outcomes.is_empty(), + "verification outcomes should be found in any item kind" + ); +} + +#[test] +fn extract_context_finds_permission_decisions() { + let items = vec![tool_exchange_item( + "Permission denied for writing to /etc/hosts", + )]; + let ctx = extract_context(&items); + assert!( + !ctx.permission_decisions.is_empty(), + "should find permission decisions" + ); +} + +#[test] +fn format_extracted_context_empty_produces_empty_string() { + let ctx = ExtractedContext::default(); + let formatted = format_extracted_context(&ctx); + assert!(formatted.is_empty()); +} + +#[test] +fn format_extracted_context_includes_all_sections() { + let ctx = ExtractedContext { + files_touched: vec!["src/main.rs".to_string()], + verification_outcomes: vec!["cargo test passed".to_string()], + permission_decisions: vec!["write permission denied".to_string()], + }; + let formatted = format_extracted_context(&ctx); + assert!(formatted.contains("FILES TOUCHED")); + assert!(formatted.contains("src/main.rs")); + assert!(formatted.contains("VERIFICATION OUTCOMES")); + assert!(formatted.contains("cargo test passed")); + assert!(formatted.contains("PERMISSION DECISIONS")); + assert!(formatted.contains("write permission denied")); +} + +#[test] +fn approx_token_count_uses_larger_of_two_heuristics() { + // Short words: "a b c d" = 4 words * 1.3 = 5.2 -> 6, chars = 7 / 4 = 2 + assert!(approx_token_count("a b c d") >= 6); + + // Long word: "abcdefghijklmnop" = 1 word * 1.3 = 2, chars = 16 / 4 = 4 + assert!(approx_token_count("abcdefghijklmnop") >= 4); +} + +#[test] +fn approx_token_count_empty_string() { + assert_eq!(approx_token_count(""), 0); +} + +#[test] +fn approx_token_count_items_sums_correctly() { + let items = vec![ + user_turn_item("hello world"), + tool_exchange_item("some tool output"), + ]; + let total = approx_token_count_items(&items); + let expected = approx_token_count("hello world") + approx_token_count("some tool output"); + assert_eq!(total, expected); +} + +// ------------------------------------------------------------------- +// M5: metadata-preserving compaction (mentra ADR-0001 §6) +// ------------------------------------------------------------------- + +fn delegation_result_item(label: &str) -> TranscriptItem { + TranscriptItem::delegation_result( + Message::user(ContentBlock::text(format!("{label} done"))), + DelegationArtifact { + kind: DelegationKind::Subagent, + agent_id: format!("agent-{label}"), + agent_name: label.to_string(), + role: None, + status: DelegationStatus::Finished, + task_summary: format!("{label} task"), + result_summary: None, + artifacts: Vec::new(), + }, + None, + ) +} + +fn with_marker(item: TranscriptItem, key: &str, value: Value) -> TranscriptItem { + item.with_details(BTreeMap::from([(key.to_string(), value)])) +} + +// Regression test 1/2: proves `select_recent_user_turns` copies its +// selections verbatim rather than rebuilding them from `Message`. A +// regression that swapped `item.clone()` for something like +// `TranscriptItem::user_turn(item.message.clone().unwrap())` would +// produce items with `details: None` here, and the derived `PartialEq` +// (which compares every field, `details` included) would catch it. +#[test] +fn select_recent_user_turns_copies_items_verbatim_details_included() { + let older = with_marker(user_turn_item("older"), "older", json!({ "keep": "older" })); + let newer = with_marker(user_turn_item("newer"), "newer", json!({ "keep": "newer" })); + let items = vec![ + older.clone(), + tool_exchange_item("not a user turn"), + newer.clone(), + ]; + + let selected = select_recent_user_turns(&items, 20_000); + + assert_eq!(selected, vec![older, newer]); +} + +// Regression test 2/2: same property for +// `select_recent_delegation_results`. +#[test] +fn select_recent_delegation_results_copies_items_verbatim_details_included() { + let first = with_marker(delegation_result_item("first"), "first", json!({ "n": 1 })); + let second = with_marker( + delegation_result_item("second"), + "second", + json!({ "n": 2 }), + ); + let items = vec![ + first.clone(), + user_turn_item("not a delegation result"), + second.clone(), + ]; + + let selected = select_recent_delegation_results(&items, 8); + + assert_eq!(selected, vec![first, second]); +} + +#[tokio::test] +async fn persist_transcript_snapshot_carries_every_items_details_bit_for_bit() { + let items = vec![ + with_marker(user_turn_item("kept"), "kept", json!({ "n": 1 })), + with_marker( + tool_exchange_item("about to be discarded"), + "about-to-be-discarded", + json!({ "n": 2 }), + ), + TranscriptItem::assistant_turn(Message::assistant(ContentBlock::text("no details here"))), + ]; + let dir = temp_dir("persist-transcript-details"); + + let path = persist_transcript(&items, &dir) + .await + .expect("persist snapshot"); + + let content = tokio::fs::read_to_string(&path) + .await + .expect("read snapshot"); + let reloaded: Vec = content + .lines() + .map(|line| serde_json::from_str(line).expect("valid TranscriptItem json")) + .collect(); + assert_eq!( + reloaded, items, + "the pre-compaction snapshot must carry every item's details bit-for-bit, \ + including items about to be discarded by summarization" + ); +} + +/// Minimal provider that returns one fixed local-summarization response — +/// enough to drive `StandardCompactionEngine::compact` end to end +/// without pulling in the full scripted-provider harness from +/// `agent::tests::support`, which is `pub(super)`-scoped to +/// `agent::tests` and unreachable from this module. +struct FixedSummaryProvider { + model: ModelInfo, +} + +#[async_trait] +impl Provider for FixedSummaryProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.model.provider.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(vec![self.model.clone()]) + } + + async fn stream(&self, _request: Request<'_>) -> Result { + Ok(provider_event_stream_from_response(Response { + id: "fixed-summary-response".to_string(), + model: self.model.id.clone(), + role: Role::Assistant, + content: vec![ContentBlock::text("test summary")], + stop_reason: None, + usage: None, + })) + } +} + +#[tokio::test] +async fn compact_preserves_salvaged_details_and_lets_discarded_details_go_with_their_items() { + let model = ModelInfo::new("test-model", "test-provider"); + let provider: Arc = Arc::new(FixedSummaryProvider { + model: model.clone(), + }); + + // Compacted-away prefix: one user turn and one delegation result + // that the engine salvages (and must copy verbatim, details + // included), plus one assistant turn and one tool exchange that are + // *not* salvaged and are honestly discarded along with their + // details. + let salvaged_user = with_marker( + user_turn_item("first message"), + "salvaged-user", + json!({ "keep": "u0" }), + ); + let discarded_assistant = + TranscriptItem::assistant_turn(Message::assistant(ContentBlock::text("ack"))); + let salvaged_delegation = with_marker( + delegation_result_item("helper"), + "salvaged-delegation", + json!({ "keep": "d0" }), + ); + let discarded_tool_result = with_marker( + tool_exchange_item("stale tool output"), + "discarded-tool", + json!({ "drop": "t0" }), + ); + // Continuation tail: kept untouched outside the compacted prefix + // (`required_tail_start_for_continuation` keeps the final + // assistant tool_use + tool result pair intact). + let tail_assistant = + TranscriptItem::assistant_turn(Message::assistant(ContentBlock::ToolUse { + id: "tail-1".to_string(), + name: "tail_tool".to_string(), + input: json!({}), + })); + let tail_result = with_marker( + TranscriptItem::tool_exchange( + Message::user(ContentBlock::text("tail tool output")), + Some("tail-1".to_string()), + false, + ), + "tail-tool", + json!({ "keep": "tail" }), + ); + + let items = vec![ + salvaged_user.clone(), + discarded_assistant.clone(), + salvaged_delegation.clone(), + discarded_tool_result.clone(), + tail_assistant.clone(), + tail_result.clone(), + ]; + let transcript = AgentTranscript::new(items.clone()); + let transcript_dir = temp_dir("m5-compaction-salvage"); + + let request = CompactionRequest { + model: model.id.clone(), + transcript, + transcript_dir, + summary_max_input_chars: 100_000, + summary_max_output_tokens: 512, + preserve_recent_user_tokens: 20_000, + preserve_recent_delegation_results: 8, + provider_request_options: ProviderRequestOptions::default(), + mode: CompactionMode::LocalOnly, + max_persisted_transcripts: None, + }; + + let outcome = StandardCompactionEngine + .compact(provider, request) + .await + .expect("compaction should not error") + .expect("compaction should produce an outcome"); + + // Counts stay consistent with the documented semantics: the whole + // compacted prefix (4 items) counts as replaced even though two of + // its items are also salvaged; only the untouched tail (2 items) + // counts as preserved_items. + assert_eq!( + outcome.replaced_items, 4, + "the whole compacted prefix counts as replaced, salvaged items included" + ); + assert_eq!( + outcome.preserved_items, 2, + "preserved_items counts only the untouched continuation tail" + ); + assert_eq!(outcome.preserved_user_turns, 1); + assert_eq!(outcome.preserved_delegation_results, 1); + + let replacement = outcome.transcript.items(); + + // Salvaged items survive verbatim, details included. + let replayed_user = replacement + .iter() + .find(|item| item.is_real_user_turn()) + .expect("salvaged user turn present in the replacement transcript"); + assert_same_entry( + replayed_user, + &salvaged_user, + "the salvaged user turn must survive bit-for-bit, details included", + ); + + let replayed_delegation = replacement + .iter() + .find(|item| item.is_delegation_result()) + .expect("salvaged delegation result present in the replacement transcript"); + assert_same_entry( + replayed_delegation, + &salvaged_delegation, + "the salvaged delegation result must survive bit-for-bit, details included", + ); + + // The untouched tail survives verbatim too. + assert_same_entry( + replacement.last().expect("a tail item"), + &tail_result, + "the untouched tail item must survive bit-for-bit, details included", + ); + + // Discarded items are honestly gone: their details never resurface + // on any other item in the replacement transcript. + let replacement_json = + serde_json::to_string(&replacement).expect("serialize replacement transcript"); + assert!( + !replacement_json.contains("discarded-tool"), + "a discarded item's details must not leak into the replacement transcript, got: {replacement_json}" + ); + + // But the pre-compaction snapshot on disk still has everything, + // including the discarded item's details — the recovery artifact + // for the summarized prefix. + let snapshot = tokio::fs::read_to_string(&outcome.transcript_path) + .await + .expect("read pre-compaction snapshot"); + let snapshot_items: Vec = snapshot + .lines() + .map(|line| serde_json::from_str(line).expect("valid TranscriptItem json")) + .collect(); + // Compared against the linked transcript rather than the raw vec the test + // assembled: appending to a transcript is what establishes an entry's + // parent, so only the linked form is what was ever snapshotted. + let linked = AgentTranscript::new(items.clone()); + assert_eq!( + snapshot_items, + linked.items(), + "the pre-compaction snapshot must preserve every original item bit-for-bit, \ + including ones the compaction goes on to discard" + ); +} + +static NEXT_TEST_DIR_ID: AtomicU64 = AtomicU64::new(1); + +fn temp_dir(label: &str) -> PathBuf { + let unique = NEXT_TEST_DIR_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time should be after unix epoch") + .as_nanos(); + std::env::temp_dir().join(format!( + "mentra-compaction-test-{label}-{timestamp}-{unique}" + )) +} + +/// Builds a request around `items` with generous, non-interfering budgets. +fn request_for(items: Vec, label: &str, model: &ModelInfo) -> CompactionRequest { + CompactionRequest { + model: model.id.clone(), + transcript: AgentTranscript::new(items), + transcript_dir: temp_dir(label), + summary_max_input_chars: 100_000, + summary_max_output_tokens: 512, + preserve_recent_user_tokens: 20_000, + preserve_recent_delegation_results: 8, + provider_request_options: ProviderRequestOptions::default(), + mode: CompactionMode::LocalOnly, + max_persisted_transcripts: None, + } +} + +fn fixed_provider(model: &ModelInfo) -> Arc { + Arc::new(FixedSummaryProvider { + model: model.clone(), + }) +} + +#[tokio::test] +async fn a_turn_pinned_by_a_tool_result_is_summarized_instead_of_refused() { + let model = ModelInfo::new("test-model", "test-provider"); + + // The whole transcript is one assistant tool call and its result: the + // continuation rule pins both, so there is nothing older to compact. + // This used to return `Ok(None)`, leaving an over-budget turn stuck. + let items = vec![ + TranscriptItem::assistant_turn(Message::assistant(ContentBlock::ToolUse { + id: "only-1".to_string(), + name: "huge_tool".to_string(), + input: json!({}), + })), + TranscriptItem::tool_exchange( + Message::user(ContentBlock::text("an enormous tool result")), + Some("only-1".to_string()), + false, + ), + ]; + + let outcome = StandardCompactionEngine + .compact( + fixed_provider(&model), + request_for(items, "m5-split-turn", &model), + ) + .await + .expect("compaction should not error") + .expect("an over-budget turn must still compact"); + + assert_eq!(outcome.replaced_items, 2, "both halves of the pair go"); + assert_eq!(outcome.preserved_items, 0, "nothing is pinned any more"); + + // Crucially, the tool result is not left without its call. + let kinds: Vec<&TranscriptKind> = outcome + .transcript + .items() + .iter() + .map(|item| &item.kind) + .collect(); + assert!( + !kinds + .iter() + .any(|kind| matches!(kind, TranscriptKind::ToolExchange { .. })), + "a tool result must never survive without the call that produced it" + ); + assert!( + kinds + .iter() + .any(|kind| matches!(kind, TranscriptKind::CompactionSummary { .. })), + "the pair is replaced by a summary" + ); +} + +#[tokio::test] +async fn a_lone_user_turn_is_never_replaced_by_a_summary_of_itself() { + let model = ModelInfo::new("test-model", "test-provider"); + let items = vec![user_turn_item("please do the thing")]; + + let outcome = StandardCompactionEngine + .compact( + fixed_provider(&model), + request_for(items, "m5-lone-user", &model), + ) + .await + .expect("compaction should not error"); + + assert!( + outcome.is_none(), + "summarizing the user's only instruction would discard the very thing \ + the turn exists to convey" + ); +} + +#[tokio::test] +async fn files_touched_accumulate_across_successive_compactions() { + let model = ModelInfo::new("test-model", "test-provider"); + + // A previous compaction already recorded a file that no surviving tool + // exchange mentions any more. + let earlier = CompactionSummary { + files_touched: vec!["src/old.rs".to_string()], + ..CompactionSummary::default() + }; + + let items = vec![ + TranscriptItem::compaction_summary(earlier), + tool_exchange_item("edited src/new.rs just now"), + TranscriptItem::assistant_turn(Message::assistant(ContentBlock::text("ack"))), + user_turn_item("carry on"), + ]; + + let outcome = StandardCompactionEngine + .compact( + fixed_provider(&model), + request_for(items, "m5-cumulative-files", &model), + ) + .await + .expect("compaction should not error") + .expect("an outcome"); + + assert!( + outcome + .summary + .files_touched + .contains(&"src/old.rs".to_string()), + "a file recorded by an earlier compaction must survive the next one; \ + got {:?}", + outcome.summary.files_touched + ); + assert!( + outcome + .summary + .files_touched + .contains(&"src/new.rs".to_string()), + "newly touched files must be added; got {:?}", + outcome.summary.files_touched + ); +} + +#[test] +fn accumulating_files_keeps_first_seen_order_and_drops_repeats() { + let merged = accumulate_files( + vec!["a.rs".to_string(), "b.rs".to_string()], + ["b.rs", "c.rs"].into_iter(), + ); + + assert_eq!(merged, vec!["a.rs", "b.rs", "c.rs"]); +} + +#[test] +fn carried_files_reads_the_newest_summary_only() { + let older = CompactionSummary { + files_touched: vec!["stale.rs".to_string()], + ..CompactionSummary::default() + }; + let newer = CompactionSummary { + files_touched: vec!["fresh.rs".to_string()], + ..CompactionSummary::default() + }; + let items = vec![ + TranscriptItem::compaction_summary(older), + TranscriptItem::compaction_summary(newer), + ]; + + // The newest summary is already cumulative, so reading only it is + // sufficient — and reading every summary would resurrect files an + // earlier round deliberately carried forward or dropped. + assert_eq!(carried_files(&items), vec!["fresh.rs".to_string()]); +} diff --git a/vendor/mentra/src/default_paths.rs b/vendor/mentra/src/default_paths.rs new file mode 100644 index 0000000..eba50e5 --- /dev/null +++ b/vendor/mentra/src/default_paths.rs @@ -0,0 +1,143 @@ +use std::{ + collections::hash_map::DefaultHasher, + hash::{Hash, Hasher}, + path::{Path, PathBuf}, +}; + +#[cfg(not(test))] +use directories::BaseDirs; + +const APP_DIR_NAME: &str = "mentra"; +const WORKSPACES_DIR_NAME: &str = "workspaces"; +const TEAM_DIR_NAME: &str = "team"; +const TASKS_DIR_NAME: &str = "tasks"; +const TRANSCRIPTS_DIR_NAME: &str = "transcripts"; +const FALLBACK_DIR_NAME: &str = ".mentra"; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct WorkspaceDefaultPaths { + pub(crate) root_dir: PathBuf, + pub(crate) default_store_path: PathBuf, + pub(crate) team_dir: PathBuf, + pub(crate) tasks_dir: PathBuf, + pub(crate) transcripts_dir: PathBuf, +} + +#[cfg(not(test))] +pub(crate) fn workspace_default_paths() -> WorkspaceDefaultPaths { + workspace_default_paths_for(canonical_workspace_dir(), platform_data_local_dir()) +} + +pub(crate) fn workspace_default_paths_for( + workspace_dir: PathBuf, + data_local_dir: Option, +) -> WorkspaceDefaultPaths { + let workspace_dir = canonicalize_or_original(workspace_dir); + let workspace_hash = workspace_hash(&workspace_dir); + let root_dir = match data_local_dir { + Some(data_local_dir) => data_local_dir + .join(APP_DIR_NAME) + .join(WORKSPACES_DIR_NAME) + .join(workspace_hash), + None => workspace_dir + .join(FALLBACK_DIR_NAME) + .join(WORKSPACES_DIR_NAME) + .join(workspace_hash), + }; + + WorkspaceDefaultPaths { + default_store_path: root_dir.join("runtime.sqlite"), + team_dir: root_dir.join(TEAM_DIR_NAME), + tasks_dir: root_dir.join(TASKS_DIR_NAME), + transcripts_dir: root_dir.join(TRANSCRIPTS_DIR_NAME), + root_dir, + } +} + +#[cfg(not(test))] +fn platform_data_local_dir() -> Option { + BaseDirs::new().map(|dirs| dirs.data_local_dir().to_path_buf()) +} + +#[cfg(not(test))] +fn canonical_workspace_dir() -> PathBuf { + canonicalize_or_original(std::env::current_dir().unwrap_or_else(|_| PathBuf::from("."))) +} + +fn canonicalize_or_original(path: PathBuf) -> PathBuf { + path.canonicalize().unwrap_or(path) +} + +fn workspace_hash(path: &Path) -> String { + let mut hasher = DefaultHasher::new(); + path.hash(&mut hasher); + format!("{:016x}", hasher.finish()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn test_path(label: &str) -> PathBuf { + std::env::temp_dir() + .join("mentra-default-paths-tests") + .join(label) + } + + #[test] + fn uses_platform_data_directory_when_available() { + let workspace = test_path("release-check-workspace"); + let data_dir = test_path("release-check-data"); + + let paths = workspace_default_paths_for(workspace.clone(), Some(data_dir.clone())); + + assert!( + paths + .root_dir + .starts_with(data_dir.join(APP_DIR_NAME).join(WORKSPACES_DIR_NAME)) + ); + assert!(paths.root_dir.ends_with(workspace_hash(&workspace))); + assert_eq!( + paths.default_store_path, + paths.root_dir.join("runtime.sqlite") + ); + assert_eq!(paths.team_dir, paths.root_dir.join(TEAM_DIR_NAME)); + assert_eq!(paths.tasks_dir, paths.root_dir.join(TASKS_DIR_NAME)); + assert_eq!( + paths.transcripts_dir, + paths.root_dir.join(TRANSCRIPTS_DIR_NAME) + ); + } + + #[test] + fn falls_back_to_workspace_dot_directory_without_platform_data_dir() { + let workspace = test_path("fallback-check-workspace"); + + let paths = workspace_default_paths_for(workspace.clone(), None); + + assert_eq!( + paths.root_dir, + workspace + .join(FALLBACK_DIR_NAME) + .join(WORKSPACES_DIR_NAME) + .join(workspace_hash(&workspace)) + ); + } + + #[test] + fn same_workspace_produces_shared_root_for_all_default_paths() { + let workspace = test_path("shared-root-workspace"); + let data_dir = test_path("shared-root-data"); + + let paths = workspace_default_paths_for(workspace, Some(data_dir)); + + for derived_path in [ + &paths.default_store_path, + &paths.team_dir, + &paths.tasks_dir, + &paths.transcripts_dir, + ] { + assert!(derived_path.starts_with(&paths.root_dir)); + } + } +} diff --git a/vendor/mentra/src/lib.rs b/vendor/mentra/src/lib.rs new file mode 100644 index 0000000..54221d0 --- /dev/null +++ b/vendor/mentra/src/lib.rs @@ -0,0 +1,84 @@ +#![doc = include_str!("../README.md")] + +mod default_paths; + +pub use mentra_provider as provider_core; + +/// Agent configuration, lifecycle, and event handling. +pub mod agent; +/// Optional OAuth helpers for provider authentication. +#[cfg(feature = "openai-oauth")] +pub mod auth; +/// Background task coordination types and services. +pub mod background; +/// Transcript compaction engine and related types. +pub mod compaction; +/// Model Context Protocol (MCP) client and tool bridge. +pub mod mcp; +/// Working-memory journal and long-term memory services. +pub mod memory; +/// Provider integrations and transport-neutral request/response types. +pub mod provider; +/// Runtime orchestration, persistence, policies, and agent APIs. +pub mod runtime; +/// Session types, metadata, and event stream primitives. +pub mod session; +/// Team coordination types and collaboration services. +pub mod team; +/// Optional test helpers for deterministic scripted runtimes. +#[cfg(any(test, feature = "test-utils"))] +pub mod test; +/// Tool traits, metadata, and builtin tools. +pub mod tool; +/// Canonical runtime transcript primitives. +pub mod transcript; + +pub use mentra_provider::{ + AnthropicRequestOptions, BuiltinProvider, ContentBlock, ContentBlockDelta, ContentBlockStart, + GeminiRequestOptions, ImageSource, Message, ModelInfo, ModelSelector, OpenAIRequestOptions, + ProviderCapabilities, ProviderCredentials, ProviderDefinition, ProviderDescriptor, + ProviderError, ProviderEvent, ProviderEventStream, ProviderId, ProviderRequestOptions, + ReasoningEffort, ReasoningFormat, ReasoningOptions, ReasoningProvenance, Request, + ResponsesRequestOptions, ResponsesStateMode, ResponsesTransport, RetryPolicy, Role, TokenUsage, + ToolChoice, ToolSearchMode, WireApi, collect_response_from_stream, + provider_event_stream_from_response, +}; + +pub use provider::{Provider, ProviderRegistry}; + +pub use agent::{ + Agent, AgentConfig, AgentWaitFuture, AgentWaitHandle, FinalOutput, QueueMode, ReasoningChange, + RoundAdjustment, RoundBoundary, RoundContext, RoundDecision, RoundStrategy, RoundToolResult, + SpawnedAgentStatus, SpawnedAgentSummary, SteeringHandle, TerminalOutputSpec, + ToolResultPagingConfig, +}; +pub use background::{BackgroundNotification, BackgroundTaskStatus, BackgroundTaskSummary}; +pub use compaction::{CompactionEngine, CompactionMode, StandardCompactionEngine}; +pub use mcp::{ + McpClientError, McpManager, McpServerConfig, McpServerStatus, McpServerSummary, McpSseClient, + McpSseConfigError, McpSseError, McpSseLimits, McpSseServerConfig, +}; +pub use runtime::{ + AgentStore, AuditStore, HybridRuntimeStore, LeaseStore, NewTask, PermissionRuleStore, RunStore, + Runtime, RuntimeBuilder, RuntimePolicy, ShellValidationMode, SkillInfo, SkillLoadError, + TaskBoard, TaskBoardError, TaskPatch, TaskStore, +}; +pub use session::{ + PermissionDecision, PermissionRequest, RememberedRule, RuleKey, RuleStore, Session, + SessionEvent, SessionEventReceiver, SessionId, SessionMetadata, SessionPermissionHandle, + SessionStatus, SubagentHandle, +}; +pub use team::{ + TeamDispatch, TeamMemberStatus, TeamMemberSummary, TeamMessage, TeamMessageKind, + TeamProtocolRequestSummary, TeamProtocolStatus, +}; +pub use tool::FileToolProfile; +pub use transcript::{ + AgentTranscript, BranchError, CompactionSummary, DelegationArtifact, DelegationEdge, + DelegationKind, DelegationStatus, EntryId, TranscriptItem, TranscriptKind, +}; + +pub mod error { + pub use crate::provider::ProviderError; + pub use crate::runtime::{ErrorCategory, RuntimeError}; +} diff --git a/vendor/mentra/src/mcp.rs b/vendor/mentra/src/mcp.rs new file mode 100644 index 0000000..34b62d8 --- /dev/null +++ b/vendor/mentra/src/mcp.rs @@ -0,0 +1,53 @@ +//! Model Context Protocol (MCP) client support. +//! +//! This module provides generic MCP clients that connect to external MCP +//! servers, discover their tools, and bridge those tools into the Mentra +//! runtime tool system. +//! +//! # Transports +//! +//! Two transports are supported, chosen by which configuration type you use: +//! +//! - **stdio** — [`McpServerConfig`] spawns a child process and speaks JSON-RPC +//! over its standard input and output. +//! - **legacy HTTP+SSE** — [`McpSseServerConfig`] opens a long-lived +//! `text/event-stream` `GET` and posts JSON-RPC messages to a second URL that +//! the server names. This is the transport from protocol revision +//! 2024-11-05, not Streamable HTTP; see [`McpSseClient`] for the distinction. +//! +//! # Architecture +//! +//! These links use absolute paths because the module's documentation is merged +//! with the outer comment on its `pub mod` declaration, which resolves relative +//! links against the crate root rather than this module. +//! +//! - [`protocol`](crate::mcp::protocol) — JSON-RPC 2.0 and MCP protocol types +//! shared by both transports +//! - [`client`](crate::mcp::client) — stdio transport client for a single MCP +//! server process +//! - [`sse`](crate::mcp::sse) — legacy HTTP+SSE transport client +//! - [`bridge`](crate::mcp::bridge) — wraps MCP tools as Mentra +//! [`ExecutableTool`] instances +//! - [`manager`](crate::mcp::manager) — manages multiple MCP server connections +//! and lifecycle +//! +//! [`ExecutableTool`]: crate::tool::ExecutableTool + +pub mod bridge; +pub mod client; +pub mod manager; +pub mod protocol; +pub mod sse; + +#[cfg(test)] +mod registration_tests; +#[cfg(test)] +mod tests; + +pub use bridge::{McpBridgedTool, mcp_tool_name, parse_mcp_tool_name}; +pub use client::{McpClientError, McpStdioClient}; +pub use manager::{McpManager, McpServerStatus, McpServerSummary}; +pub use protocol::{McpServerConfig, McpToolDefinition}; +pub use sse::client::{McpSseClient, McpSseError}; +pub use sse::config::{McpSseConfigError, McpSseLimits, McpSseServerConfig, SecretString}; +pub use sse::endpoint::EndpointError; diff --git a/vendor/mentra/src/mcp/bridge.rs b/vendor/mentra/src/mcp/bridge.rs new file mode 100644 index 0000000..2c6fba4 --- /dev/null +++ b/vendor/mentra/src/mcp/bridge.rs @@ -0,0 +1,200 @@ +//! Bridge that wraps MCP server tools as Mentra `ExecutableTool` instances. + +use std::sync::Arc; + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::tool::{ + ParallelToolContext, RuntimeToolDescriptor, ToolApprovalCategory, ToolCapability, + ToolDefinition, ToolDurability, ToolExecutionCategory, ToolExecutor, ToolResult, + ToolSideEffectLevel, +}; + +use super::client::McpStdioClient; +use super::protocol::{McpToolCallResult, McpToolDefinition}; +use super::sse::client::McpSseClient; + +/// The transport-independent surface [`McpBridgedTool`] needs from a client. +/// +/// Each transport reports failures with its own error type, so this trait +/// flattens them to a message rather than forcing a shared error enum on the +/// public clients. +/// +/// This is a sealed trait: it is public only so that +/// [`McpBridgedTool::new`] can be generic over the transport, and it is not +/// implementable outside this crate. +#[async_trait] +pub trait McpToolClient: sealed::Sealed + Send + Sync { + /// Calls one tool, rendering any transport failure as a message. + async fn call_tool( + &self, + tool_name: &str, + arguments: Option, + ) -> Result; +} + +mod sealed { + /// Prevents outside implementations of [`super::McpToolClient`]. + pub trait Sealed {} + + impl Sealed for super::McpStdioClient {} + impl Sealed for super::McpSseClient {} + + #[cfg(test)] + impl Sealed for crate::mcp::tests::SuccessfulMcpClient {} +} + +#[async_trait] +impl McpToolClient for McpStdioClient { + async fn call_tool( + &self, + tool_name: &str, + arguments: Option, + ) -> Result { + McpStdioClient::call_tool(self, tool_name, arguments) + .await + .map_err(|error| error.to_string()) + } +} + +#[async_trait] +impl McpToolClient for McpSseClient { + async fn call_tool( + &self, + tool_name: &str, + arguments: Option, + ) -> Result { + McpSseClient::call_tool(self, tool_name, arguments) + .await + .map_err(|error| error.to_string()) + } +} + +/// Prefix applied to MCP tool names to namespace them. +const MCP_TOOL_PREFIX: &str = "mcp__"; + +/// Construct the namespaced tool name for an MCP tool. +pub fn mcp_tool_name(server_name: &str, tool_name: &str) -> String { + format!("{MCP_TOOL_PREFIX}{server_name}__{tool_name}") +} + +/// Parse a namespaced MCP tool name back into `(server_name, tool_name)`. +pub fn parse_mcp_tool_name(name: &str) -> Option<(&str, &str)> { + let rest = name.strip_prefix(MCP_TOOL_PREFIX)?; + let (server, tool) = rest.split_once("__")?; + Some((server, tool)) +} + +/// A Mentra tool backed by an MCP server tool. +pub struct McpBridgedTool { + server_name: String, + tool_def: McpToolDefinition, + client: Arc, +} + +impl McpBridgedTool { + /// Wraps one tool from a connected MCP server. + /// + /// The client is generic over the transport, so this accepts an + /// `Arc` and an `Arc` alike. + pub fn new(server_name: String, tool_def: McpToolDefinition, client: Arc) -> Self + where + C: McpToolClient + 'static, + { + Self::from_client(server_name, tool_def, client) + } + + fn from_client( + server_name: String, + tool_def: McpToolDefinition, + client: Arc, + ) -> Self { + Self { + server_name, + tool_def, + client, + } + } + + #[cfg(test)] + pub(crate) fn new_for_test( + server_name: String, + tool_def: McpToolDefinition, + client: Arc, + ) -> Self { + Self::from_client(server_name, tool_def, client) + } + + fn full_name(&self) -> String { + mcp_tool_name(&self.server_name, &self.tool_def.name) + } +} + +impl std::fmt::Debug for McpBridgedTool { + /// Renders the bridged identity without reaching into the client, which + /// holds transport credentials. + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("McpBridgedTool") + .field("name", &self.full_name()) + .finish_non_exhaustive() + } +} + +impl ToolDefinition for McpBridgedTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + let description = self.tool_def.description.clone().unwrap_or_default(); + + let input_schema = self + .tool_def + .input_schema + .clone() + .unwrap_or_else(|| json!({"type": "object", "properties": {}})); + + RuntimeToolDescriptor::builder(self.full_name()) + .description(description) + .input_schema(input_schema) + .capability(ToolCapability::Custom(format!("mcp:{}", self.server_name))) + .side_effect_level(ToolSideEffectLevel::External) + .durability(ToolDurability::Ephemeral) + .execution_category(ToolExecutionCategory::ExclusiveLocalMutation) + .approval_category(ToolApprovalCategory::Process) + .build() + } +} + +#[async_trait] +impl ToolExecutor for McpBridgedTool { + async fn execute(&self, _ctx: ParallelToolContext, input: Value) -> ToolResult { + let arguments = if input.is_null() + || (input.is_object() && input.as_object().is_none_or(|o| o.is_empty())) + { + None + } else { + Some(input) + }; + + let result = self + .client + .call_tool(&self.tool_def.name, arguments) + .await + .map_err(|error| format!("MCP tool call failed: {error}"))?; + + // Concatenate text content blocks into the result string. + let mut output = String::new(); + for block in &result.content { + if let Some(text) = &block.text { + if !output.is_empty() { + output.push('\n'); + } + output.push_str(text); + } + } + + if result.is_error { + Err(output) + } else { + Ok(output) + } + } +} diff --git a/vendor/mentra/src/mcp/client.rs b/vendor/mentra/src/mcp/client.rs new file mode 100644 index 0000000..b0a221e --- /dev/null +++ b/vendor/mentra/src/mcp/client.rs @@ -0,0 +1,354 @@ +//! MCP stdio client — spawns a child process and communicates via JSON-RPC over stdin/stdout. + +#[cfg(test)] +mod tests; + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::Duration; + +use serde::de::DeserializeOwned; +use serde_json::Value as JsonValue; +use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; +use tokio::process::{Child, ChildStdin, Command}; +use tokio::sync::{Mutex, oneshot}; + +use super::protocol::*; + +/// Default timeout for the MCP `initialize` handshake. +const INITIALIZE_TIMEOUT: Duration = Duration::from_secs(10); +/// Default timeout for `tools/list`. +const LIST_TOOLS_TIMEOUT: Duration = Duration::from_secs(30); +/// Default timeout for `tools/call`. +const CALL_TOOL_TIMEOUT: Duration = Duration::from_secs(120); +/// Bound on how many `tools/list` pages are followed. +/// +/// Cursors are opaque, so a server repeating one cannot be detected by value; +/// only a page bound stops the walk. +const MAX_TOOL_PAGES: usize = 1_000; + +/// Errors from the MCP stdio client. +#[derive(Debug, thiserror::Error)] +pub enum McpClientError { + #[error("failed to spawn MCP server process: {0}")] + SpawnFailed(#[from] std::io::Error), + + #[error("MCP server process has no stdin")] + NoStdin, + + #[error("MCP server process has no stdout")] + NoStdout, + + #[error("MCP server returned JSON-RPC error: {0}")] + JsonRpc(JsonRpcError), + + #[error("timeout waiting for MCP response ({0:?})")] + Timeout(Duration), + + #[error("MCP server process exited unexpectedly")] + ProcessExited, + + #[error("failed to parse MCP response: {0}")] + ParseError(String), + + #[error("MCP server kept paginating tools/list past {limit} pages")] + TooManyToolPages { limit: usize }, + + #[error("MCP client is already shut down")] + Shutdown, +} + +type PendingMap = HashMap>>; + +/// A running MCP stdio client connected to one server process. +pub struct McpStdioClient { + stdin: Mutex, + _child: Mutex, + next_id: AtomicU64, + pending: Arc>, + server_info: Option, + tools: Vec, + server_name: String, +} + +impl McpStdioClient { + /// Spawn the MCP server process and perform the `initialize` handshake. + pub async fn connect(config: &McpServerConfig) -> Result { + let mut cmd = Command::new(&config.command); + cmd.args(&config.args) + .stdin(std::process::Stdio::piped()) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::null()); + + for (key, value) in &config.env { + cmd.env(key, value); + } + if let Some(cwd) = &config.cwd { + cmd.current_dir(cwd); + } + + let mut child = cmd.spawn()?; + + let stdin = child.stdin.take().ok_or(McpClientError::NoStdin)?; + let stdout = child.stdout.take().ok_or(McpClientError::NoStdout)?; + + let pending: Arc> = Arc::new(Mutex::new(HashMap::new())); + + // Spawn the reader task that routes responses to pending callers. + let pending_clone = pending.clone(); + tokio::spawn(async move { + let mut reader = BufReader::new(stdout); + let mut line = String::new(); + loop { + line.clear(); + match reader.read_line(&mut line).await { + Ok(0) | Err(_) => break, + Ok(_) => {} + } + let trimmed = line.trim(); + if trimmed.is_empty() { + continue; + } + if let Ok(resp) = serde_json::from_str::(trimmed) { + // A response carries a result or an error. A server-initiated + // request such as `ping` also has an id, and without this + // check it would resolve the caller holding that id with a + // null result. + if resp.result.is_none() && resp.error.is_none() { + continue; + } + let id = match &resp.id { + JsonRpcId::Number(n) => *n, + _ => continue, + }; + let mut pending = pending_clone.lock().await; + if let Some(tx) = pending.remove(&id) { + let result = if let Some(err) = resp.error { + Err(McpClientError::JsonRpc(err)) + } else { + Ok(resp.result.unwrap_or(JsonValue::Null)) + }; + let _ = tx.send(result); + } + } + } + // When the reader exits, signal all pending callers. + let mut pending = pending_clone.lock().await; + for (_, tx) in pending.drain() { + let _ = tx.send(Err(McpClientError::ProcessExited)); + } + }); + + let mut client = Self { + stdin: Mutex::new(stdin), + _child: Mutex::new(child), + next_id: AtomicU64::new(1), + pending, + server_info: None, + tools: Vec::new(), + server_name: config.name.clone(), + }; + + // Perform initialize handshake. + client.initialize().await?; + + // Discover tools. + client.discover_tools().await?; + + Ok(client) + } + + /// Server name from the configuration. + pub fn server_name(&self) -> &str { + &self.server_name + } + + /// Server info returned by the `initialize` handshake. + pub fn server_info(&self) -> Option<&McpServerInfo> { + self.server_info.as_ref() + } + + /// Tools discovered from this server. + pub fn tools(&self) -> &[McpToolDefinition] { + &self.tools + } + + /// Send a JSON-RPC request and wait for the response. + async fn call( + &self, + method: &str, + params: Option

, + timeout_duration: Duration, + ) -> Result { + let id = self.next_id.fetch_add(1, Ordering::Relaxed); + + let params_value = params + .map(|p| serde_json::to_value(p).expect("serialize params")) + .filter(|v| !v.is_null()); + + let request = JsonRpcRequest::new(id, method, params_value); + let mut line = serde_json::to_string(&request).expect("serialize request"); + line.push('\n'); + + let (tx, rx) = oneshot::channel(); + { + let mut pending = self.pending.lock().await; + pending.insert(id, tx); + } + + { + let mut stdin = self.stdin.lock().await; + if stdin.write_all(line.as_bytes()).await.is_err() || stdin.flush().await.is_err() { + // The request never reached the server, so drop its + // registration rather than leaving it to time out. + self.pending.lock().await.remove(&id); + return Err(McpClientError::ProcessExited); + } + } + + let result = match tokio::time::timeout(timeout_duration, rx).await { + Ok(Ok(result)) => result?, + Ok(Err(_)) => return Err(McpClientError::ProcessExited), + Err(_) => { + // Remove the registration so a timed-out request cannot leak an + // entry for the lifetime of the connection. + self.pending.lock().await.remove(&id); + return Err(McpClientError::Timeout(timeout_duration)); + } + }; + + serde_json::from_value(result) + .map_err(|e| McpClientError::ParseError(format!("deserialize response: {e}"))) + } + + /// Send a JSON-RPC notification (no response expected). + async fn notify( + &self, + method: &str, + params: Option

, + ) -> Result<(), McpClientError> { + // Notifications have no id — use a raw object. + let mut obj = serde_json::json!({ + "jsonrpc": "2.0", + "method": method, + }); + if let Some(p) = params { + obj["params"] = serde_json::to_value(p).expect("serialize params"); + } + let mut line = serde_json::to_string(&obj).expect("serialize notification"); + line.push('\n'); + + let mut stdin = self.stdin.lock().await; + stdin + .write_all(line.as_bytes()) + .await + .map_err(|_| McpClientError::ProcessExited)?; + stdin + .flush() + .await + .map_err(|_| McpClientError::ProcessExited)?; + Ok(()) + } + + async fn initialize(&mut self) -> Result<(), McpClientError> { + let params = McpInitializeParams { + protocol_version: "2024-11-05".to_string(), + capabilities: serde_json::json!({}), + client_info: McpClientInfo { + name: "mentra".to_string(), + version: env!("CARGO_PKG_VERSION").to_string(), + }, + }; + + let result: McpInitializeResult = self + .call("initialize", Some(params), INITIALIZE_TIMEOUT) + .await?; + + self.server_info = Some(result.server_info); + + // Send initialized notification. + self.notify::("notifications/initialized", None) + .await?; + + Ok(()) + } + + async fn discover_tools(&mut self) -> Result<(), McpClientError> { + let mut all_tools = Vec::new(); + let mut cursor: Option = None; + let mut pages = 0_usize; + + loop { + let params = McpListToolsParams { + cursor: cursor.clone(), + }; + let result: McpListToolsResult = self + .call("tools/list", Some(params), LIST_TOOLS_TIMEOUT) + .await?; + + all_tools.extend(result.tools); + + pages += 1; + if pages >= MAX_TOOL_PAGES { + // A server that keeps handing back a cursor would otherwise + // loop forever, growing the tool list without bound. Cursors + // are opaque, so a repeat cannot be detected by value. + return Err(McpClientError::TooManyToolPages { + limit: MAX_TOOL_PAGES, + }); + } + + match result.next_cursor { + Some(next) if !next.is_empty() => cursor = Some(next), + _ => break, + } + } + + self.tools = all_tools; + Ok(()) + } + + /// Call a tool on this server. + pub async fn call_tool( + &self, + tool_name: &str, + arguments: Option, + ) -> Result { + self.call_tool_with_timeout(tool_name, arguments, CALL_TOOL_TIMEOUT) + .await + } + + /// Call a tool on this server, bounding the wait explicitly. + pub async fn call_tool_with_timeout( + &self, + tool_name: &str, + arguments: Option, + timeout: Duration, + ) -> Result { + let params = McpToolCallParams { + name: tool_name.to_string(), + arguments, + }; + self.call("tools/call", Some(params), timeout).await + } + + /// The number of requests still awaiting a response. + #[cfg(test)] + pub(crate) async fn pending_len(&self) -> usize { + self.pending.lock().await.len() + } + + /// Shut down the MCP server process gracefully. + pub async fn shutdown(&self) { + // Best-effort: drop stdin to signal the child. + let mut stdin = self.stdin.lock().await; + drop(stdin.shutdown().await); + } +} + +impl Drop for McpStdioClient { + fn drop(&mut self) { + // The child process will be killed when the Child handle is dropped. + } +} diff --git a/vendor/mentra/src/mcp/client/tests.rs b/vendor/mentra/src/mcp/client/tests.rs new file mode 100644 index 0000000..77c6a11 --- /dev/null +++ b/vendor/mentra/src/mcp/client/tests.rs @@ -0,0 +1,197 @@ +//! Tests for the MCP stdio client, driven by a scripted server process. +//! +//! The server is a short Python program so the test can control exactly which +//! JSON-RPC frames come back, including ones a well-behaved server would never +//! send. Tests that need it are skipped when no interpreter is available rather +//! than failing, so the suite still runs on a machine without Python. + +use std::collections::HashMap; +use std::time::Duration; + +use super::{McpClientError, McpStdioClient}; +use crate::mcp::protocol::McpServerConfig; + +/// Returns an interpreter that can run the scripted server, if one exists. +fn python() -> Option<&'static str> { + ["python3", "python"].into_iter().find(|candidate| { + std::process::Command::new(candidate) + .arg("--version") + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .status() + .is_ok_and(|status| status.success()) + }) +} + +/// Builds a config running the given Python source as an MCP server. +fn scripted_server(python: &str, source: &str) -> McpServerConfig { + McpServerConfig { + name: "scripted".to_string(), + command: python.to_string(), + args: vec!["-c".to_string(), source.to_string()], + env: HashMap::new(), + cwd: None, + } +} + +/// A server that completes the handshake, then behaves as `extra` directs. +/// +/// `extra` runs after `tools/list`, receiving each subsequent request line. +fn handshake_server(extra: &str) -> String { + format!( + r#" +import sys, json + +def send(payload): + sys.stdout.write(json.dumps(payload) + "\n") + sys.stdout.flush() + +def read(): + line = sys.stdin.readline() + if not line: + raise SystemExit(0) + return json.loads(line) + +# initialize +request = read() +send({{"jsonrpc": "2.0", "id": request["id"], "result": {{ + "protocolVersion": "2024-11-05", + "capabilities": {{}}, + "serverInfo": {{"name": "scripted", "version": "9.9.9"}}}}}}) + +# notifications/initialized carries no id and expects no reply +read() + +# tools/list +request = read() +send({{"jsonrpc": "2.0", "id": request["id"], "result": {{"tools": [ + {{"name": "echo", "inputSchema": {{"type": "object"}}}}]}}}}) + +{extra} +"# + ) +} + +#[tokio::test] +async fn completes_the_handshake_and_discovers_tools() { + let Some(python) = python() else { + eprintln!("skipping: no Python interpreter available"); + return; + }; + + let config = scripted_server(python, &handshake_server("read()")); + let client = McpStdioClient::connect(&config) + .await + .expect("the handshake should succeed"); + + assert_eq!( + client.server_info().map(|info| info.name.as_str()), + Some("scripted") + ); + assert_eq!(client.tools().len(), 1); + assert_eq!(client.tools()[0].name, "echo"); +} + +/// A server-initiated request carries a method and an id but no result. Without +/// a guard the reader treats it as a response and resolves whichever caller +/// happens to hold that id with a null result. +#[tokio::test] +async fn a_server_initiated_request_does_not_resolve_a_pending_call() { + let Some(python) = python() else { + eprintln!("skipping: no Python interpreter available"); + return; + }; + + let extra = r#" +# tools/call — answer with a ping request first, reusing the caller's id. +request = read() +send({"jsonrpc": "2.0", "id": request["id"], "method": "ping"}) +send({"jsonrpc": "2.0", "id": request["id"], "result": { + "content": [{"type": "text", "text": "real result"}], "isError": False}}) +read() +"#; + + let config = scripted_server(python, &handshake_server(extra)); + let client = McpStdioClient::connect(&config) + .await + .expect("the handshake should succeed"); + + let result = client + .call_tool("echo", None) + .await + .expect("the real response should resolve the call"); + + assert_eq!( + result.content[0].text.as_deref(), + Some("real result"), + "a ping request must not be mistaken for the response" + ); +} + +/// A request that times out must remove its pending entry. A leak here is +/// bounded by request count, but a long-lived agent session makes many. +#[tokio::test] +async fn a_timed_out_request_does_not_leak_its_pending_entry() { + let Some(python) = python() else { + eprintln!("skipping: no Python interpreter available"); + return; + }; + + let extra = r#" +# Swallow one tools/call without answering, then answer the next. +read() +request = read() +send({"jsonrpc": "2.0", "id": request["id"], "result": { + "content": [{"type": "text", "text": "second call"}], "isError": False}}) +read() +"#; + + let config = scripted_server(python, &handshake_server(extra)); + let client = McpStdioClient::connect(&config) + .await + .expect("the handshake should succeed"); + + let timed_out = tokio::time::timeout( + Duration::from_secs(5), + client.call_tool_with_timeout("echo", None, Duration::from_millis(150)), + ) + .await + .expect("the call should give up on its own") + .expect_err("no response arrives for the first call"); + assert!(matches!(timed_out, McpClientError::Timeout(_))); + + assert_eq!( + client.pending_len().await, + 0, + "a timed-out request must not leave an entry behind" + ); + + let result = client + .call_tool("echo", None) + .await + .expect("the connection should remain usable"); + assert_eq!(result.content[0].text.as_deref(), Some("second call")); +} + +#[tokio::test] +async fn every_pending_call_fails_when_the_process_exits() { + let Some(python) = python() else { + eprintln!("skipping: no Python interpreter available"); + return; + }; + + // Exit immediately after the handshake, without answering the tool call. + let config = scripted_server(python, &handshake_server("raise SystemExit(0)")); + let client = McpStdioClient::connect(&config) + .await + .expect("the handshake should succeed"); + + let error = tokio::time::timeout(Duration::from_secs(10), client.call_tool("echo", None)) + .await + .expect("the call must fail rather than hang") + .expect_err("a dead process cannot answer"); + assert!( + matches!(error, McpClientError::ProcessExited), + "got {error:?}" + ); +} diff --git a/vendor/mentra/src/mcp/manager.rs b/vendor/mentra/src/mcp/manager.rs new file mode 100644 index 0000000..b3bf8e5 --- /dev/null +++ b/vendor/mentra/src/mcp/manager.rs @@ -0,0 +1,262 @@ +//! Manages multiple MCP server connections and their lifecycle. + +use std::collections::HashMap; +use std::sync::Arc; + +use super::bridge::{McpBridgedTool, McpToolClient, mcp_tool_name}; +use super::client::{McpClientError, McpStdioClient}; +use super::protocol::{McpServerConfig, McpToolDefinition}; +use super::sse::client::{McpSseClient, McpSseError}; +use super::sse::config::McpSseServerConfig; + +/// Status of an MCP server connection. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum McpServerStatus { + Disconnected, + Connecting, + Connected, + Error, +} + +impl std::fmt::Display for McpServerStatus { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Disconnected => write!(f, "disconnected"), + Self::Connecting => write!(f, "connecting"), + Self::Connected => write!(f, "connected"), + Self::Error => write!(f, "error"), + } + } +} + +/// Summary of a managed MCP server. +#[derive(Debug, Clone)] +pub struct McpServerSummary { + pub name: String, + pub status: McpServerStatus, + pub server_version: Option, + pub tool_count: usize, + pub error: Option, +} + +/// A connected client, whichever transport it speaks. +/// +/// The manager needs more than [`McpToolClient`] provides — it reports server +/// versions and shuts connections down — so the transports are held in an enum +/// rather than behind that trait. +enum TransportClient { + Stdio(Arc), + Sse(Arc), +} + +impl TransportClient { + /// The server version reported by the `initialize` handshake. + fn server_version(&self) -> Option { + match self { + Self::Stdio(client) => client.server_info().map(|info| info.version.clone()), + Self::Sse(client) => client.server_info().map(|info| info.version.clone()), + } + } + + /// Closes the connection. + async fn shutdown(&self) { + match self { + Self::Stdio(client) => client.shutdown().await, + Self::Sse(client) => client.shutdown().await, + } + } + + /// Calls a tool, flattening the transport's error to a message. + async fn call_tool( + &self, + tool_name: &str, + arguments: Option, + ) -> Result { + match self { + Self::Stdio(client) => McpToolClient::call_tool(&**client, tool_name, arguments).await, + Self::Sse(client) => McpToolClient::call_tool(&**client, tool_name, arguments).await, + } + } + + /// Bridges every advertised tool into a runtime tool. + fn bridge(&self, server_name: &str, tools: &[McpToolDefinition]) -> Vec { + tools + .iter() + .map(|tool| match self { + Self::Stdio(client) => { + McpBridgedTool::new(server_name.to_string(), tool.clone(), client.clone()) + } + Self::Sse(client) => { + McpBridgedTool::new(server_name.to_string(), tool.clone(), client.clone()) + } + }) + .collect() + } +} + +/// Tracks a connected MCP server. +struct ConnectedServer { + client: TransportClient, + tools: Vec, +} + +/// Manages the lifecycle of multiple MCP server processes. +pub struct McpManager { + servers: HashMap, + errors: HashMap, +} + +impl McpManager { + pub fn new() -> Self { + Self { + servers: HashMap::new(), + errors: HashMap::new(), + } + } + + /// Connect to an MCP server over stdio and discover its tools. + /// Returns the bridged tools ready for registration. + pub async fn connect( + &mut self, + config: &McpServerConfig, + ) -> Result, McpClientError> { + // Disconnect existing connection if any. + self.disconnect(&config.name).await; + + let client = McpStdioClient::connect(config).await.inspect_err(|e| { + self.errors.insert(config.name.clone(), e.to_string()); + })?; + + let tools = client.tools().to_vec(); + let client = TransportClient::Stdio(Arc::new(client)); + + Ok(self.register(config.name.clone(), client, tools)) + } + + /// Connect to an MCP server over the legacy HTTP+SSE transport and discover + /// its tools. + /// + /// Returns the bridged tools ready for registration, exactly as + /// [`connect`](Self::connect) does for stdio. + pub async fn connect_sse( + &mut self, + config: &McpSseServerConfig, + ) -> Result, McpSseError> { + self.disconnect(&config.name).await; + + let client = McpSseClient::connect(config).await.inspect_err(|error| { + self.errors.insert(config.name.clone(), error.to_string()); + })?; + + let tools = client.tools().to_vec(); + let client = TransportClient::Sse(Arc::new(client)); + + Ok(self.register(config.name.clone(), client, tools)) + } + + /// Records a connected server and bridges its tools. + fn register( + &mut self, + name: String, + client: TransportClient, + tools: Vec, + ) -> Vec { + let bridged = client.bridge(&name, &tools); + self.errors.remove(&name); + self.servers.insert(name, ConnectedServer { client, tools }); + bridged + } + + /// Disconnect a server by name. + pub async fn disconnect(&mut self, name: &str) { + if let Some(server) = self.servers.remove(name) { + server.client.shutdown().await; + } + } + + /// Shut down all connected servers. + pub async fn shutdown_all(&mut self) { + let names: Vec = self.servers.keys().cloned().collect(); + for name in names { + self.disconnect(&name).await; + } + } + + /// List all server summaries. + pub fn list_servers(&self) -> Vec { + let mut summaries: Vec = self + .servers + .iter() + .map(|(name, server)| McpServerSummary { + name: name.clone(), + status: McpServerStatus::Connected, + server_version: server.client.server_version(), + tool_count: server.tools.len(), + error: None, + }) + .collect(); + + // Include errored servers. + for (name, error) in &self.errors { + if !self.servers.contains_key(name) { + summaries.push(McpServerSummary { + name: name.clone(), + status: McpServerStatus::Error, + server_version: None, + tool_count: 0, + error: Some(error.clone()), + }); + } + } + + summaries.sort_by(|a, b| a.name.cmp(&b.name)); + summaries + } + + /// Get the namespaced tool names for all connected servers. + pub fn all_tool_names(&self) -> Vec { + self.servers + .iter() + .flat_map(|(name, server)| { + server + .tools + .iter() + .map(move |tool| mcp_tool_name(name, &tool.name)) + }) + .collect() + } + + /// Call a tool on a specific server, whichever transport it speaks. + /// + /// Each transport reports failures with its own error type, so this returns + /// the message rather than widening the error into a shared enum. + pub async fn call_tool( + &self, + server_name: &str, + tool_name: &str, + arguments: Option, + ) -> Result { + let server = self + .servers + .get(server_name) + .ok_or_else(|| format!("MCP server '{server_name}' not connected"))?; + + server.client.call_tool(tool_name, arguments).await + } + + /// Check if a server is connected. + pub fn is_connected(&self, name: &str) -> bool { + self.servers.contains_key(name) + } + + /// Number of connected servers. + pub fn connected_count(&self) -> usize { + self.servers.len() + } +} + +impl Default for McpManager { + fn default() -> Self { + Self::new() + } +} diff --git a/vendor/mentra/src/mcp/protocol.rs b/vendor/mentra/src/mcp/protocol.rs new file mode 100644 index 0000000..48aaf2d --- /dev/null +++ b/vendor/mentra/src/mcp/protocol.rs @@ -0,0 +1,181 @@ +//! JSON-RPC 2.0 and MCP protocol types. + +use serde::{Deserialize, Serialize}; +use serde_json::Value as JsonValue; + +// --------------------------------------------------------------------------- +// JSON-RPC 2.0 primitives +// --------------------------------------------------------------------------- + +/// JSON-RPC 2.0 request identifier. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum JsonRpcId { + Number(u64), + String(String), + Null, +} + +/// Outbound JSON-RPC 2.0 request. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct JsonRpcRequest { + pub jsonrpc: String, + pub id: JsonRpcId, + pub method: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub params: Option, +} + +impl JsonRpcRequest { + pub fn new(id: u64, method: impl Into, params: Option) -> Self { + Self { + jsonrpc: "2.0".to_string(), + id: JsonRpcId::Number(id), + method: method.into(), + params, + } + } +} + +/// JSON-RPC 2.0 error object. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct JsonRpcError { + pub code: i64, + pub message: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub data: Option, +} + +impl std::fmt::Display for JsonRpcError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "JSON-RPC error {}: {}", self.code, self.message) + } +} + +/// Inbound JSON-RPC 2.0 response. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct JsonRpcResponse { + pub jsonrpc: String, + pub id: JsonRpcId, + #[serde(skip_serializing_if = "Option::is_none")] + pub result: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub error: Option, +} + +// --------------------------------------------------------------------------- +// MCP initialize +// --------------------------------------------------------------------------- + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct McpInitializeParams { + pub protocol_version: String, + pub capabilities: JsonValue, + pub client_info: McpClientInfo, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpClientInfo { + pub name: String, + pub version: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct McpInitializeResult { + pub protocol_version: String, + pub capabilities: JsonValue, + pub server_info: McpServerInfo, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpServerInfo { + pub name: String, + #[serde(default)] + pub version: String, +} + +// --------------------------------------------------------------------------- +// MCP tools/list +// --------------------------------------------------------------------------- + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct McpListToolsParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub cursor: Option, +} + +/// A tool definition returned by the MCP server. +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct McpToolDefinition { + pub name: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub input_schema: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct McpListToolsResult { + pub tools: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub next_cursor: Option, +} + +// --------------------------------------------------------------------------- +// MCP tools/call +// --------------------------------------------------------------------------- + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpToolCallParams { + pub name: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub arguments: Option, +} + +/// A single content block in a tool call result. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpToolCallContent { + #[serde(rename = "type")] + pub kind: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub text: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub data: Option, + #[serde(rename = "mimeType", skip_serializing_if = "Option::is_none")] + pub mime_type: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct McpToolCallResult { + pub content: Vec, + #[serde(default)] + pub is_error: bool, +} + +// --------------------------------------------------------------------------- +// MCP server configuration +// --------------------------------------------------------------------------- + +/// Configuration for a single MCP server. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct McpServerConfig { + /// Display name for the server. + pub name: String, + /// The command to spawn. + pub command: String, + /// Arguments to pass to the command. + #[serde(default)] + pub args: Vec, + /// Extra environment variables. + #[serde(default)] + pub env: std::collections::HashMap, + /// Working directory for the spawned process. + #[serde(skip_serializing_if = "Option::is_none")] + pub cwd: Option, +} diff --git a/vendor/mentra/src/mcp/registration_tests.rs b/vendor/mentra/src/mcp/registration_tests.rs new file mode 100644 index 0000000..943079a --- /dev/null +++ b/vendor/mentra/src/mcp/registration_tests.rs @@ -0,0 +1,572 @@ +//! Tests for MCP registration through the runtime builder. + +use serde_json::json; + +use crate::mcp::sse::testing::SseTestServer; +use crate::mcp::{McpManager, McpSseServerConfig, mcp_tool_name}; + +const REMOTE_CANARY: &str = "REMOTE_CANARY_MUST_NOT_SURFACE"; + +/// Scripts a fixture through the handshake, advertising the given tools. +/// +/// Returns the manager alongside the bridged tools: the manager owns the +/// connection those tools call through, so dropping it would close the stream. +async fn connect_sse( + server: &SseTestServer, + tools: serde_json::Value, +) -> (Vec, McpManager) { + let config = McpSseServerConfig::new("obs", server.sse_url()); + let connecting = tokio::spawn(async move { + let mut manager = McpManager::new(); + let bridged = manager.connect_sse(&config).await; + (bridged, manager) + }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=abc"); + + server.wait_for_posts(1); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 1, + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "fixture", "version": "4.5.6"} + } + })); + + server.wait_for_posts(3); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 2, + "result": {"tools": tools} + })); + + let (bridged, manager) = connecting.await.expect("no panic"); + (bridged.expect("the handshake should succeed"), manager) +} + +#[tokio::test(flavor = "multi_thread")] +async fn the_manager_bridges_sse_tools_under_a_namespaced_name() { + let server = SseTestServer::start(); + let (bridged, _manager) = connect_sse( + &server, + json!([ + {"name": "search_logs", "description": "Search logs", "inputSchema": {"type": "object"}}, + {"name": "list_alerts", "inputSchema": {"type": "object"}} + ]), + ) + .await; + + use crate::tool::ToolDefinition; + let names: Vec = bridged + .iter() + .map(|tool| tool.descriptor().name.to_string()) + .collect(); + + assert_eq!( + names, + vec![ + mcp_tool_name("obs", "search_logs"), + mcp_tool_name("obs", "list_alerts"), + ], + "SSE tools must be namespaced exactly like stdio tools" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn a_bridged_sse_tool_carries_its_description_and_schema() { + let server = SseTestServer::start(); + let (bridged, _manager) = connect_sse( + &server, + json!([{ + "name": "search_logs", + "description": "Search the log corpus", + "inputSchema": { + "type": "object", + "properties": {"query": {"type": "string"}}, + "required": ["query"] + } + }]), + ) + .await; + + use crate::tool::ToolDefinition; + let descriptor = bridged[0].descriptor(); + assert_eq!( + descriptor.description.as_deref(), + Some("Search the log corpus") + ); + assert_eq!( + descriptor.input_schema["properties"]["query"]["type"], "string", + "the server's schema must reach the model unchanged" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn the_manager_reports_a_connected_sse_server() { + let server = SseTestServer::start(); + let config = McpSseServerConfig::new("obs", server.sse_url()); + let connecting = tokio::spawn(async move { + let mut manager = McpManager::new(); + let result = manager.connect_sse(&config).await; + (result.map(|tools| tools.len()), manager) + }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=abc"); + server.wait_for_posts(1); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 1, + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {}, + "serverInfo": {"name": "fixture", "version": "4.5.6"} + } + })); + server.wait_for_posts(3); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 2, + "result": {"tools": [{"name": "search", "inputSchema": {"type": "object"}}]} + })); + + let (count, manager) = connecting.await.expect("no panic"); + assert_eq!(count.expect("the handshake should succeed"), 1); + + assert!(manager.is_connected("obs")); + assert_eq!(manager.connected_count(), 1); + assert_eq!( + manager.all_tool_names(), + vec![mcp_tool_name("obs", "search")] + ); + + let summary = manager + .list_servers() + .into_iter() + .find(|summary| summary.name == "obs") + .expect("the server should be listed"); + assert_eq!(summary.status, crate::mcp::McpServerStatus::Connected); + assert_eq!(summary.server_version.as_deref(), Some("4.5.6")); + assert_eq!(summary.tool_count, 1); +} + +#[tokio::test(flavor = "multi_thread")] +async fn a_failed_sse_connection_is_recorded_as_an_error() { + let server = SseTestServer::with_opening(crate::mcp::sse::testing::StreamOpening::Status { + code: 404, + body: "no such stream".to_string(), + }); + + let mut manager = McpManager::new(); + let config = McpSseServerConfig::new("obs", server.sse_url()); + manager + .connect_sse(&config) + .await + .expect_err("a 404 must not connect"); + + assert!(!manager.is_connected("obs")); + let summary = manager + .list_servers() + .into_iter() + .find(|summary| summary.name == "obs") + .expect("an errored server should still be listed"); + assert_eq!(summary.status, crate::mcp::McpServerStatus::Error); + assert!(summary.error.is_some()); +} + +#[tokio::test(flavor = "multi_thread")] +async fn manager_error_summaries_do_not_retain_json_rpc_text() { + let server = SseTestServer::start(); + let config = McpSseServerConfig::new("obs", server.sse_url()); + let connecting = tokio::spawn(async move { + let mut manager = McpManager::new(); + let error = manager + .connect_sse(&config) + .await + .expect_err("the initialize error must fail the connection"); + (error, manager) + }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=abc"); + server.wait_for_posts(1); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 1, + "error": { + "code": -32001, + "message": REMOTE_CANARY, + "data": {"forged": REMOTE_CANARY} + } + })); + + let (error, manager) = connecting.await.expect("no panic"); + assert!( + matches!(error, crate::mcp::McpSseError::JsonRpc(ref rpc) if rpc.code == -32001), + "got {error:?}" + ); + + let summary = manager + .list_servers() + .into_iter() + .find(|summary| summary.name == "obs") + .expect("the failed server should be listed"); + let rendered = format!("{summary:?}"); + assert!(!rendered.contains(REMOTE_CANARY), "got {rendered}"); + assert!( + summary + .error + .as_deref() + .is_some_and(|message| message.contains("-32001")), + "the safe summary should preserve the JSON-RPC code: {rendered}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn a_misconfigured_sse_server_fails_before_any_connection_is_opened() { + let mut manager = McpManager::new(); + // A token over plaintext to a non-loopback host is refused at validation. + let config = + McpSseServerConfig::new("obs", "http://internal.corp/sse").with_bearer_token("secret"); + + let error = manager + .connect_sse(&config) + .await + .expect_err("validation must reject this before dialing"); + assert!( + !error.to_string().contains("secret"), + "the error must not echo the credential: {error}" + ); +} + +/// `build_async` must actually register the MCP tools it connects, and a +/// runtime built without any MCP server must not gain namespaced tools. +/// +/// Without this, disabling the whole registration arm in `build_async` leaves +/// the suite green: the runtime still builds, it just silently advertises +/// nothing. That failure is invisible until an agent cannot find its tools. +#[tokio::test(flavor = "multi_thread")] +async fn build_async_registers_the_tools_of_a_connected_sse_server() { + use crate::Runtime; + + let server = SseTestServer::start(); + let config = McpSseServerConfig::new("obs", server.sse_url()); + + let building = tokio::spawn(async move { + Runtime::empty_builder() + .with_provider_instance(StubProvider) + .with_mcp_sse_server(config) + .build_async() + .await + }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=abc"); + server.wait_for_posts(1); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 1, + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {}, + "serverInfo": {"name": "fixture", "version": "4.5.6"} + } + })); + server.wait_for_posts(3); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 2, + "result": {"tools": [ + {"name": "search_logs", "inputSchema": {"type": "object"}}, + {"name": "list_alerts", "inputSchema": {"type": "object"}} + ]} + })); + + let runtime = building + .await + .expect("no panic") + .expect("the runtime should build"); + + let registered: Vec = runtime + .tools() + .into_iter() + .map(|tool| tool.name.to_string()) + .filter(|name| name.starts_with("mcp__")) + .collect(); + + assert_eq!( + registered.len(), + 2, + "build_async must register every tool the server advertised, got {registered:?}" + ); + assert!(registered.contains(&mcp_tool_name("obs", "search_logs"))); + assert!(registered.contains(&mcp_tool_name("obs", "list_alerts"))); + + assert!( + runtime + .tool_descriptor(&mcp_tool_name("obs", "search_logs")) + .is_some(), + "a registered MCP tool must be resolvable by name" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn build_async_registers_no_mcp_tools_without_a_configured_server() { + use crate::Runtime; + + let runtime = Runtime::empty_builder() + .with_provider_instance(StubProvider) + .build_async() + .await + .expect("the runtime should build"); + + let registered: Vec = runtime + .tools() + .into_iter() + .map(|tool| tool.name.to_string()) + .filter(|name| name.starts_with("mcp__")) + .collect(); + + assert!( + registered.is_empty(), + "no MCP server was configured, got {registered:?}" + ); +} + +/// A provider that satisfies the builder's "at least one provider" check. +/// +/// The registration tests never send a request, so it only needs to exist. +#[derive(Clone)] +struct StubProvider; + +#[async_trait::async_trait] +impl crate::provider::Provider for StubProvider { + fn descriptor(&self) -> crate::provider::ProviderDescriptor { + crate::provider::ProviderDescriptor::new(crate::BuiltinProvider::Anthropic) + } + + async fn list_models(&self) -> Result, crate::provider::ProviderError> { + Ok(vec![crate::ModelInfo::new( + "stub-model", + crate::BuiltinProvider::Anthropic, + )]) + } + + async fn stream( + &self, + _request: crate::provider::Request<'_>, + ) -> Result { + let (_tx, rx) = tokio::sync::mpsc::unbounded_channel(); + Ok(rx) + } +} + +/// An SSE-backed tool must reach the model through exactly the same result +/// limiter and paging path as a stdio tool or a custom tool. This runs a real +/// fixture server behind a bridged tool inside a scripted runtime and compares +/// its transcript entry byte for byte against a custom tool returning the same +/// text. +#[tokio::test(flavor = "multi_thread")] +async fn sse_tool_output_is_limited_exactly_like_a_custom_tool() { + use std::collections::BTreeMap; + + use crate::{ + ContentBlock, + runtime::RuntimePolicy, + test::{MockRuntime, MockToolCall}, + tool::ToolDefinition, + }; + + let full_output = "one\ntwo\nthree"; + let server = SseTestServer::start(); + let (bridged, _manager) = connect_sse( + &server, + json!([{"name": "large_output", "inputSchema": {"type": "object"}}]), + ) + .await; + let bridged_name = bridged[0].descriptor().name.to_string(); + + let mock = MockRuntime::builder() + .with_policy( + RuntimePolicy::permissive() + .with_max_tool_result_bytes(8) + .with_max_tool_result_lines(1) + .spill_full_tool_output(false), + ) + .tool_calls([ + MockToolCall::new(&bridged_name, json!({})).with_id("sse-call"), + MockToolCall::new("matching_custom_output", json!({})).with_id("custom-call"), + ]) + .text("done") + .build() + .expect("build mock runtime"); + + for tool in bridged { + mock.runtime().register_tool(tool); + } + mock.runtime().register_tool(EchoTool { + output: full_output.to_string(), + }); + + let mut agent = mock + .runtime() + .spawn("mcp-sse-truncation-test", mock.model()) + .expect("spawn agent"); + + let running = tokio::spawn(async move { + let response = agent + .send(vec![ContentBlock::text("run both tools")]) + .await + .expect("run agent"); + (response, agent) + }); + + // Answer the bridged tool call once the runtime has posted it. + server.wait_for_posts(4); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": { + "content": [{"type": "text", "text": full_output}], + "isError": false + } + })); + + let (response, _agent) = running.await.expect("no panic"); + assert_eq!(response.text(), "done"); + + let requests = mock.recorded_requests().await; + let provider_results = requests[1] + .messages + .iter() + .flat_map(|message| &message.content) + .filter_map(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => Some((tool_use_id.as_str(), (content.as_str(), *is_error))), + _ => None, + }) + .collect::>(); + + let sse_result = provider_results + .get("sse-call") + .expect("the SSE tool result should reach the provider"); + let custom_result = provider_results + .get("custom-call") + .expect("the custom tool result should reach the provider"); + + assert_eq!( + sse_result, custom_result, + "an SSE tool must be limited identically to any other tool" + ); + assert_eq!( + *sse_result, + ( + "one\n[truncated: showing 1 of 3 lines; full output was not saved because spill-to-file is disabled by runtime policy]", + false, + ) + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn bridged_sse_errors_do_not_put_json_rpc_text_in_model_context() { + use crate::{ + ContentBlock, + test::{MockRuntime, MockToolCall}, + tool::ToolDefinition, + }; + + let server = SseTestServer::start(); + let (bridged, _manager) = connect_sse( + &server, + json!([{"name": "fail", "inputSchema": {"type": "object"}}]), + ) + .await; + let bridged_name = bridged[0].descriptor().name.to_string(); + + let mock = MockRuntime::builder() + .tool_calls([MockToolCall::new(&bridged_name, json!({})).with_id("sse-error")]) + .text("done") + .build() + .expect("build mock runtime"); + for tool in bridged { + mock.runtime().register_tool(tool); + } + + let mut agent = mock + .runtime() + .spawn("mcp-sse-error-redaction-test", mock.model()) + .expect("spawn agent"); + let running = tokio::spawn(async move { + agent + .send(vec![ContentBlock::text("call the failing tool")]) + .await + }); + + server.wait_for_posts(4); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "error": { + "code": -32602, + "message": REMOTE_CANARY, + "data": {"forged": REMOTE_CANARY} + } + })); + + let response = running + .await + .expect("no panic") + .expect("the agent should continue after the tool error"); + assert_eq!(response.text(), "done"); + + let requests = mock.recorded_requests().await; + let (content, is_error) = requests[1] + .messages + .iter() + .flat_map(|message| &message.content) + .find_map(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } if tool_use_id == "sse-error" => Some((content.as_str(), *is_error)), + _ => None, + }) + .expect("the provider should receive the bridged error"); + + assert!(is_error); + assert!(!content.contains(REMOTE_CANARY), "got {content}"); + assert!(content.contains("-32602"), "got {content}"); +} + +/// A custom tool returning a fixed string, used as the limiter baseline. +struct EchoTool { + output: String, +} + +impl crate::tool::ToolDefinition for EchoTool { + fn descriptor(&self) -> crate::tool::ToolSpec { + crate::tool::ToolSpec::builder("matching_custom_output") + .description("Return the same output as the MCP test tool") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(crate::tool::ToolSideEffectLevel::External) + .build() + } +} + +#[async_trait::async_trait] +impl crate::tool::ToolExecutor for EchoTool { + async fn execute( + &self, + _ctx: crate::tool::ParallelToolContext, + _input: serde_json::Value, + ) -> crate::tool::ToolResult { + Ok(self.output.clone()) + } +} diff --git a/vendor/mentra/src/mcp/sse.rs b/vendor/mentra/src/mcp/sse.rs new file mode 100644 index 0000000..2a68967 --- /dev/null +++ b/vendor/mentra/src/mcp/sse.rs @@ -0,0 +1,14 @@ +//! Legacy MCP HTTP+SSE transport (protocol revision 2024-11-05). +//! +//! This is the transport MCP defined in revision 2024-11-05, where the client +//! holds a long-lived `GET` stream open and posts JSON-RPC messages to a +//! separate URL that the server names. It is distinct from Streamable HTTP; +//! see [`client`](crate::mcp::sse::client) for the differences and the full +//! lifecycle. + +pub mod client; +pub mod config; +pub mod endpoint; +#[cfg(test)] +pub(crate) mod testing; +pub(crate) mod wire; diff --git a/vendor/mentra/src/mcp/sse/client.rs b/vendor/mentra/src/mcp/sse/client.rs new file mode 100644 index 0000000..29f27d5 --- /dev/null +++ b/vendor/mentra/src/mcp/sse/client.rs @@ -0,0 +1,790 @@ +//! MCP client for the legacy HTTP+SSE transport (protocol revision 2024-11-05). +//! +//! # The transport +//! +//! This is the *older* MCP HTTP transport, not Streamable HTTP. The two are +//! easy to confuse and are not interchangeable: +//! +//! | | legacy HTTP+SSE (this module) | Streamable HTTP | +//! |---|---|---| +//! | Endpoints | a `GET` stream plus a separate `POST` URL | one URL for both | +//! | POST target | named by the server in an `endpoint` event | the configured URL | +//! | Responses | always on the `GET` stream | in the POST response or a stream | +//! | Session | a query parameter in the endpoint URL | the `Mcp-Session-Id` header | +//! +//! Servers that answer `404` on `/mcp` but serve `/sse` require this transport. +//! +//! # Lifecycle +//! +//! 1. `GET` the configured URL with `Accept: text/event-stream`. +//! 2. Wait for an `event: endpoint` frame naming the `POST` URL, resolve it +//! against the configured URL, and require it to stay on the same origin. +//! 3. `POST` JSON-RPC requests as `application/json`. Any 2xx — including the +//! `202 Accepted` both reference servers return — means the message was +//! accepted for processing, not that it completed. +//! 4. Read JSON-RPC responses from `event: message` frames on the stream and +//! correlate them to requests by id. +//! +//! The handshake is `initialize`, then a `notifications/initialized` +//! notification, then a paginated `tools/list`. +//! +//! # Failure behavior +//! +//! The stream carries every response, so losing it ends the session. This +//! client fails closed: when the stream ends, every pending request resolves +//! with an error rather than hanging. It never reconnects and never re-sends a +//! `tools/call`, because an MCP tool may have side effects and a transparent +//! retry would execute it twice with no caller involvement. +//! +//! A `tools/call` whose `POST` may have reached the server but whose response +//! never arrived is reported as [`McpSseError::RequestIndeterminate`] rather +//! than as a plain failure, so a caller can tell "may have run" apart from +//! "definitely did not". + +#[cfg(test)] +mod tests; + +use std::collections::HashMap; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::{Arc, Mutex, MutexGuard}; +use std::time::Duration; + +use futures_util::StreamExt; +use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; +use serde::de::DeserializeOwned; +use serde_json::Value as JsonValue; +use tokio::sync::oneshot; +use tokio::task::JoinHandle; +use url::Url; + +use super::config::{McpSseConfigError, McpSseLimits, McpSseServerConfig}; +use super::endpoint::{EndpointError, resolve_endpoint}; +use super::wire::{SseParser, SseWireError}; +use crate::mcp::protocol::*; + +/// The protocol revision this transport implements. +const PROTOCOL_VERSION: &str = "2024-11-05"; + +/// Bound on how much of an HTTP error body is read for diagnostics. +/// +/// The body is attacker-controlled, so it is never included in an error; this +/// bound exists only so that reading and discarding it cannot be turned into a +/// memory-exhaustion primitive. +const MAX_DIAGNOSTIC_BODY_BYTES: usize = 8 * 1024; + +/// Errors from the MCP SSE client. +/// +/// No variant produced from a server response carries a response body, an SSE +/// payload, a server-controlled free-form metadata value, a JSON-RPC message or +/// data value, a tool argument, or a tool result. Server text is never +/// interpolated into an error, because a malicious server would otherwise be +/// able to write arbitrary content — including forged log lines and terminal +/// escape sequences — into an operator's logs or a model's context. Fixed +/// metadata such as an HTTP status or JSON-RPC code remains available. +#[derive(Debug, thiserror::Error)] +pub enum McpSseError { + #[error("invalid MCP SSE configuration: {0}")] + Config(#[from] McpSseConfigError), + + #[error("invalid MCP SSE endpoint: {0}")] + Endpoint(#[from] EndpointError), + + #[error("failed to reach the MCP SSE server: {0}")] + Transport(String), + + #[error("MCP SSE server answered the {method} request with HTTP {status}")] + HttpStatus { + method: &'static str, + status: reqwest::StatusCode, + }, + + #[error( + "MCP SSE server answered with a redirect, which is not followed because it would send \ + credentials to an unvalidated origin" + )] + RedirectRefused, + + #[error( + "MCP SSE server answered with content type '{content_type}', expected text/event-stream" + )] + UnexpectedContentType { content_type: String }, + + #[error("MCP SSE stream framing error: {0}")] + Wire(#[from] SseWireError), + + #[error("MCP SSE endpoint event exceeded the {limit} byte limit")] + EndpointTooLarge { limit: usize }, + + #[error("MCP SSE server kept paginating tools/list past {limit} pages")] + TooManyToolPages { limit: usize }, + + #[error("MCP SSE server returned JSON-RPC error: {0}")] + JsonRpc(JsonRpcError), + + #[error("failed to parse the MCP SSE response: {0}")] + ParseError(String), + + #[error("timed out after {0:?} waiting for the MCP SSE server")] + Timeout(Duration), + + #[error("the MCP SSE stream closed before the request completed")] + StreamClosed, + + /// The request may have reached the server, but no response arrived before + /// the stream ended or the deadline passed. + /// + /// The call may have executed. The `POST` and the response travel on + /// different connections, so a server can accept and run a tool while the + /// stream dies, and the client cannot tell that apart from the tool never + /// starting. Callers must not retry automatically: an MCP tool can send + /// mail, charge a card, or write a file. + /// + #[error( + "the MCP SSE server may have received the '{method}' request but never answered it; \ + the call may have executed and must not be retried automatically" + )] + RequestIndeterminate { method: String }, + + #[error("the MCP SSE client is shut down")] + Shutdown, +} + +/// The outcome of a JSON-RPC request, kept in the pending map. +type PendingReply = Result; + +/// Correlates in-flight requests to the responses arriving on the stream. +/// +/// Lookups remove the entry, so the first response for an id wins and a +/// malicious server cannot deliver a second result for a call the caller has +/// already observed. +#[derive(Default)] +struct Pending { + waiters: HashMap, + /// Set once the stream ends so late requests fail immediately rather than + /// waiting out their timeout. + closed: bool, +} + +struct PendingWaiter { + reply: oneshot::Sender, + method: String, +} + +/// Removes one pending waiter if its request future is dropped. +/// +/// Request futures are cancellation points while sending the POST, draining +/// its response body, and waiting on the SSE stream. A synchronous mutex keeps +/// this cleanup available to `Drop`, where awaiting a Tokio mutex is +/// impossible. The critical sections only mutate the in-memory map and never +/// perform I/O. +struct PendingRegistration { + pending: Arc>, + id: u64, +} + +impl Drop for PendingRegistration { + fn drop(&mut self) { + lock_pending(&self.pending).waiters.remove(&self.id); + } +} + +/// A connected MCP server speaking the legacy HTTP+SSE transport. +/// +/// This is the low-level client. It performs the handshake, exposes the +/// server's advertised tools, and calls one selected tool. It deliberately does +/// not register anything with the runtime, so a host can apply its own +/// allowlists, redaction, and evidence policy over the top. Use +/// `RuntimeBuilder::with_mcp_sse_server` when the generic bridging behavior is +/// what you want. +pub struct McpSseClient { + http: reqwest::Client, + /// The `POST` target named by the server, validated to the configured origin. + endpoint: Url, + headers: HeaderMap, + limits: McpSseLimits, + next_id: AtomicU64, + pending: Arc>, + reader: JoinHandle<()>, + server_info: Option, + tools: Vec, + server_name: String, + stream_url: Url, +} + +impl McpSseClient { + /// Opens the SSE stream, performs the MCP handshake, and discovers tools. + pub async fn connect(config: &McpSseServerConfig) -> Result { + let stream_url = config.validate()?; + let headers = build_headers(config)?; + let http = build_http_client(&config.limits)?; + + let response = tokio::time::timeout( + config.limits.connect_timeout, + http.get(stream_url.clone()) + .header(reqwest::header::ACCEPT, "text/event-stream") + .headers(headers.clone()) + .send(), + ) + .await + .map_err(|_| McpSseError::Timeout(config.limits.connect_timeout))? + .map_err(transport_error)?; + + check_stream_response(&response)?; + + let pending: Arc> = Arc::new(Mutex::new(Pending::default())); + let (endpoint_tx, endpoint_rx) = oneshot::channel(); + + let reader = tokio::spawn(read_stream( + response, + Arc::clone(&pending), + endpoint_tx, + config.limits.clone(), + )); + + // The endpoint event must arrive before anything can be sent. Bound the + // wait: a buffering proxy is a common cause of it never arriving. + // + // Every failure from here on must abort the reader before returning, or + // the task and the connection it holds outlive the failed connect. + let endpoint = match tokio::time::timeout(config.limits.connect_timeout, endpoint_rx).await + { + Ok(Ok(Ok(raw))) => match resolve_endpoint(&stream_url, &raw) { + Ok(endpoint) => endpoint, + Err(error) => { + reader.abort(); + return Err(error.into()); + } + }, + Ok(Ok(Err(error))) => { + reader.abort(); + return Err(error); + } + Ok(Err(_)) => { + reader.abort(); + return Err(McpSseError::StreamClosed); + } + Err(_) => { + reader.abort(); + return Err(McpSseError::Timeout(config.limits.connect_timeout)); + } + }; + + let mut client = Self { + http, + endpoint, + headers, + limits: config.limits.clone(), + next_id: AtomicU64::new(1), + pending, + reader, + server_info: None, + tools: Vec::new(), + server_name: config.name.clone(), + stream_url, + }; + + // A failure here returns `client` by value, so its `Drop` aborts the + // reader; there is no separate cleanup path to keep in sync. + client.initialize().await?; + client.discover_tools().await?; + + Ok(client) + } + + /// The configured name of this server, used to namespace its tools. + pub fn server_name(&self) -> &str { + &self.server_name + } + + /// The configured SSE stream URL. + pub fn stream_url(&self) -> &Url { + &self.stream_url + } + + /// Server information returned by the `initialize` handshake. + pub fn server_info(&self) -> Option<&McpServerInfo> { + self.server_info.as_ref() + } + + /// The tools this server advertised. + pub fn tools(&self) -> &[McpToolDefinition] { + &self.tools + } + + /// Calls one tool on this server. + pub async fn call_tool( + &self, + tool_name: &str, + arguments: Option, + ) -> Result { + let params = McpToolCallParams { + name: tool_name.to_string(), + arguments, + }; + self.request("tools/call", Some(params), self.limits.call_tool_timeout) + .await + } + + /// Closes the stream and fails every request still in flight. + pub async fn shutdown(&self) { + self.reader.abort(); + let mut pending = lock_pending(&self.pending); + pending.closed = true; + drain_pending(&mut pending); + } + + /// Sends a JSON-RPC request and waits for its correlated response. + async fn request( + &self, + method: &'static str, + params: Option

, + timeout: Duration, + ) -> Result { + let id = self.next_id.fetch_add(1, Ordering::Relaxed); + let params = params + .map(|params| serde_json::to_value(params)) + .transpose() + .map_err(|error| McpSseError::ParseError(error.to_string()))? + .filter(|params| !params.is_null()); + let request = JsonRpcRequest::new(id, method, params); + + let (reply_tx, reply_rx) = oneshot::channel(); + let _registration = { + // Register before sending. The server answers the POST before it + // processes the message, so the response can reach the stream + // before the POST future resolves. + let mut pending = lock_pending(&self.pending); + if pending.closed { + return Err(McpSseError::StreamClosed); + } + pending.waiters.insert( + id, + PendingWaiter { + reply: reply_tx, + method: method.to_string(), + }, + ); + PendingRegistration { + pending: Arc::clone(&self.pending), + id, + } + }; + + // One deadline covers the complete operation. In particular, a peer + // cannot evade `call_tool_timeout` by accepting the TCP connection and + // withholding either the HTTP response head or its declared body. + let operation = async { + self.post(&request) + .await + .map_err(|error| classify_post_failure(method, error))?; + + match reply_rx.await { + Ok(result) => result, + // The reader dropped the sender, which only happens on + // teardown. A tools/call may already have executed. + Err(_) => Err(indeterminate(method)), + } + }; + let result = match tokio::time::timeout(timeout, operation).await { + Ok(result) => result?, + Err(_) => return Err(request_timeout(method, timeout)), + }; + + serde_json::from_value(result) + .map_err(|_| McpSseError::ParseError("response shape did not match MCP".to_string())) + } + + /// Sends a JSON-RPC notification, which expects no response. + async fn notify(&self, method: &str, timeout: Duration) -> Result<(), McpSseError> { + let notification = serde_json::json!({"jsonrpc": "2.0", "method": method}); + tokio::time::timeout(timeout, self.post(¬ification)) + .await + .map_err(|_| McpSseError::Timeout(timeout))? + } + + /// `POST`s one JSON-RPC message to the validated endpoint. + async fn post(&self, message: &T) -> Result<(), McpSseError> { + let response = self + .http + .post(self.endpoint.clone()) + .headers(self.headers.clone()) + .json(message) + .send() + .await + .map_err(transport_error)?; + + let status = response.status(); + if status.is_redirection() { + return Err(McpSseError::RedirectRefused); + } + if !status.is_success() { + // Read a bounded prefix so the connection returns to the pool, then + // discard it: the body is attacker-controlled and never surfaces. + drain_bounded(response).await; + return Err(McpSseError::HttpStatus { + method: "POST", + status, + }); + } + + // The JSON-RPC result never arrives here — it comes back on the stream. + // Drain anyway so the connection is reusable. + drain_bounded(response).await; + Ok(()) + } + + /// Performs the `initialize` handshake and the follow-up notification. + async fn initialize(&mut self) -> Result<(), McpSseError> { + let params = McpInitializeParams { + protocol_version: PROTOCOL_VERSION.to_string(), + capabilities: serde_json::json!({}), + client_info: McpClientInfo { + name: "mentra".to_string(), + version: env!("CARGO_PKG_VERSION").to_string(), + }, + }; + + let result: McpInitializeResult = self + .request("initialize", Some(params), self.limits.initialize_timeout) + .await?; + self.server_info = Some(result.server_info); + + self.notify("notifications/initialized", self.limits.initialize_timeout) + .await + } + + /// Walks the paginated `tools/list` cursor to the end. + async fn discover_tools(&mut self) -> Result<(), McpSseError> { + let mut tools = Vec::new(); + let mut cursor: Option = None; + let mut pages = 0_usize; + + loop { + let params = McpListToolsParams { + cursor: cursor.clone(), + }; + let page: McpListToolsResult = self + .request("tools/list", Some(params), self.limits.list_tools_timeout) + .await?; + tools.extend(page.tools); + + pages += 1; + if pages >= self.limits.max_tool_pages { + // A server that keeps handing back a cursor would otherwise + // loop forever, growing the tool list without bound. The + // cursor is opaque, so a repeat cannot be detected by value. + return Err(McpSseError::TooManyToolPages { + limit: self.limits.max_tool_pages, + }); + } + + match page.next_cursor { + // A missing or empty cursor means the last page. + Some(next) if !next.is_empty() => cursor = Some(next), + _ => break, + } + } + + self.tools = tools; + Ok(()) + } +} + +impl Drop for McpSseClient { + fn drop(&mut self) { + // Cancel the reader so the task and its connection do not outlive the + // client that owns them. + self.reader.abort(); + } +} + +impl std::fmt::Debug for McpSseClient { + /// Renders without the header map, which holds credentials. + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("McpSseClient") + .field("server_name", &self.server_name) + .field("stream_url", &self.stream_url.as_str()) + .field("tools", &self.tools.len()) + .finish_non_exhaustive() + } +} + +/// Reads the SSE stream until it ends, routing frames to waiting callers. +async fn read_stream( + response: reqwest::Response, + pending: Arc>, + endpoint_tx: oneshot::Sender>, + limits: McpSseLimits, +) { + let mut parser = SseParser::new(limits.max_event_bytes); + let mut body = response.bytes_stream(); + let mut endpoint_tx = Some(endpoint_tx); + + loop { + let next = tokio::time::timeout(limits.stream_idle_timeout, body.next()).await; + + let chunk = match next { + Ok(Some(Ok(chunk))) => chunk, + // Any stream error is terminal. reqwest's `is_body` does not + // reliably identify body errors, so it is not consulted. + Ok(Some(Err(_))) | Ok(None) => break, + Err(_) => break, + }; + + let events = match parser.feed(&chunk) { + Ok(events) => events, + Err(error) => { + // Framing can no longer be trusted, so the stream is torn down + // rather than resynchronized at an attacker-chosen boundary. + notify_endpoint(&mut endpoint_tx, Err(McpSseError::Wire(error))); + break; + } + }; + + for event in events { + match event.event.as_str() { + "endpoint" => { + if event.data.len() > limits.max_endpoint_bytes { + notify_endpoint( + &mut endpoint_tx, + Err(McpSseError::EndpointTooLarge { + limit: limits.max_endpoint_bytes, + }), + ); + // Fail closed: without a usable endpoint nothing can be sent. + break; + } + // Only the first endpoint event is honored. A later one + // would silently redirect in-flight traffic. + notify_endpoint(&mut endpoint_tx, Ok(event.data)); + } + "message" => deliver_message(&pending, &event.data), + // Unknown event names, including the `ping` frames older + // sse-starlette versions emit, are ignored rather than fatal. + _ => {} + } + } + } + + // Whatever ended the stream, no further response can arrive. + notify_endpoint(&mut endpoint_tx, Err(McpSseError::StreamClosed)); + let mut pending = lock_pending(&pending); + pending.closed = true; + drain_pending(&mut pending); +} + +/// Routes one `message` frame to the caller waiting on its id. +fn deliver_message(pending: &Arc>, data: &str) { + let Ok(response) = serde_json::from_str::(data) else { + // Malformed JSON, or a server-initiated request such as `ping`. Neither + // is a response, so neither is correlated. Dropping it keeps a hostile + // server from turning unsolicited frames into per-id state. + return; + }; + + // A response carries a result or an error; anything else is a request. + if response.result.is_none() && response.error.is_none() { + return; + } + + let JsonRpcId::Number(id) = response.id else { + return; + }; + + // Removing the entry means the first response wins and a repeated id + // cannot deliver a second result for an already-observed call. + let Some(waiter) = lock_pending(pending).waiters.remove(&id) else { + return; + }; + + let reply = match response.error { + Some(error) => Err(McpSseError::JsonRpc(JsonRpcError { + code: error.code, + message: "server message omitted".to_string(), + data: None, + })), + None => Ok(response.result.unwrap_or(JsonValue::Null)), + }; + let _ = waiter.reply.send(reply); +} + +/// Sends the endpoint outcome exactly once. +fn notify_endpoint( + endpoint_tx: &mut Option>>, + outcome: Result, +) { + if let Some(tx) = endpoint_tx.take() { + let _ = tx.send(outcome); + } +} + +/// Fails every pending request. +/// +/// A request that is still registered may already have been sent. A +/// `tools/call` here may therefore have executed, and saying so is what lets a +/// caller avoid re-running a non-idempotent action. +fn drain_pending(pending: &mut Pending) { + for (_, waiter) in pending.waiters.drain() { + let _ = waiter.reply.send(Err(indeterminate(&waiter.method))); + } +} + +/// Reports a sent-but-unanswered request in the terms a caller needs. +/// +/// Only `tools/call` is reported as indeterminate: the handshake methods are +/// idempotent, so an unanswered one is simply a closed stream. +fn indeterminate(method: &str) -> McpSseError { + if method == "tools/call" { + McpSseError::RequestIndeterminate { + method: method.to_string(), + } + } else { + McpSseError::StreamClosed + } +} + +/// Converts a POST failure into the method-level certainty the caller needs. +/// +/// Once a `tools/call` POST future has begun, neither a transport error nor an +/// HTTP response proves that the application did not process its body first. +/// HTTP status describes the response, not an atomic absence of server-side +/// effects, so every POST failure is indeterminate for a potentially mutating +/// tool call. Handshake methods retain the underlying diagnostic because they +/// are safe to establish again in a new session. +fn classify_post_failure(method: &str, error: McpSseError) -> McpSseError { + if method == "tools/call" { + indeterminate(method) + } else { + error + } +} + +/// Reports expiration of the whole request operation. +fn request_timeout(method: &str, timeout: Duration) -> McpSseError { + if method == "tools/call" { + indeterminate(method) + } else { + McpSseError::Timeout(timeout) + } +} + +/// Acquires the pending map even if another task panicked while holding it. +/// +/// Losing the map on poison would strand unrelated requests forever. No map +/// mutation runs user code, so recovering the contained value is safe. +fn lock_pending(pending: &Mutex) -> MutexGuard<'_, Pending> { + pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) +} + +/// Builds the HTTP client shared by the stream and the message endpoint. +fn build_http_client(limits: &McpSseLimits) -> Result { + reqwest::Client::builder() + // Redirects are never followed. reqwest strips `Authorization` only + // across hosts and compares host and port without scheme, so an + // https->http redirect on the same host would keep the credential and + // send it in the clear. The request body is never stripped at all, so a + // followed redirect would also hand tool arguments to the new target. + .redirect(reqwest::redirect::Policy::none()) + // Reqwest retries selected protocol-level rejections on its own once + // the negotiated protocol can signal them (HTTP/2 REFUSED_STREAM and + // kin). A tools/call POST may already have executed by the time such + // a signal arrives, so an automatic resend would replay a + // side-effecting call with no caller involvement — the exact + // double-execution this client's request path refuses. Today's + // feature set negotiates HTTP/1.1 only, where no such signal exists; + // this pin keeps the no-replay guarantee structural rather than an + // accident of the current feature graph. + .retry(reqwest::retry::never()) + .connect_timeout(limits.connect_timeout) + // Deliberately no `.timeout()`: that is a total deadline covering the + // response body, which would kill the long-lived stream on a fixed + // interval. The request and notification paths apply their own total + // deadlines around each finite POST operation. + .build() + .map_err(|error| McpSseError::Transport(error.to_string())) +} + +/// Converts configured headers into a map, marking every value sensitive. +fn build_headers(config: &McpSseServerConfig) -> Result { + let mut headers = HeaderMap::new(); + + for (name, value) in &config.headers { + let name = HeaderName::try_from(name.as_str()).map_err(|_| { + McpSseConfigError::InvalidHeaderName { + name: name.to_string(), + } + })?; + let mut value = HeaderValue::try_from(value.expose_secret()).map_err(|_| { + McpSseConfigError::InvalidHeaderValue { + name: name.to_string(), + } + })?; + // A plain HeaderValue prints its contents in Debug output. Marking it + // sensitive redacts it there and tells HTTP/2 not to index it. + value.set_sensitive(true); + headers.insert(name, value); + } + + Ok(headers) +} + +/// Requires a 200 response carrying an event stream. +fn check_stream_response(response: &reqwest::Response) -> Result<(), McpSseError> { + let status = response.status(); + if status.is_redirection() { + return Err(McpSseError::RedirectRefused); + } + if !status.is_success() { + return Err(McpSseError::HttpStatus { + method: "GET", + status, + }); + } + + let content_type = response + .headers() + .get(reqwest::header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(); + + // Prefix match: servers commonly send `text/event-stream; charset=utf-8`. + if !content_type + .trim_start() + .to_ascii_lowercase() + .starts_with("text/event-stream") + { + return Err(McpSseError::UnexpectedContentType { + content_type: "[server value omitted]".to_string(), + }); + } + + Ok(()) +} + +/// Reads and discards a bounded prefix of a response body. +async fn drain_bounded(response: reqwest::Response) { + let mut body = response.bytes_stream(); + let mut seen = 0_usize; + while let Some(Ok(chunk)) = body.next().await { + seen += chunk.len(); + if seen >= MAX_DIAGNOSTIC_BODY_BYTES { + break; + } + } +} + +/// Renders a transport failure without echoing server-controlled text. +fn transport_error(error: reqwest::Error) -> McpSseError { + let reason = if error.is_timeout() { + "timed out" + } else if error.is_connect() { + "could not connect" + } else if error.is_request() { + "the request could not be sent" + } else { + "the connection failed" + }; + McpSseError::Transport(reason.to_string()) +} diff --git a/vendor/mentra/src/mcp/sse/client/tests.rs b/vendor/mentra/src/mcp/sse/client/tests.rs new file mode 100644 index 0000000..3b6f16c --- /dev/null +++ b/vendor/mentra/src/mcp/sse/client/tests.rs @@ -0,0 +1,1504 @@ +//! Tests for the legacy MCP HTTP+SSE client, driven by a local fixture. + +use serde_json::json; + +use super::{McpSseClient, McpSseError}; +use crate::mcp::sse::config::{McpSseLimits, McpSseServerConfig}; +use crate::mcp::sse::testing::{PostReply, SseTestServer, StreamOpening}; + +const REMOTE_CANARY: &str = "REMOTE_CANARY_MUST_NOT_SURFACE"; + +fn assert_remote_canary_absent(error: &McpSseError) { + let display = error.to_string(); + let debug = format!("{error:?}"); + assert!(!display.contains(REMOTE_CANARY), "got {display}"); + assert!(!debug.contains(REMOTE_CANARY), "got {debug}"); +} + +/// The `initialize` result every handshake test replies with. +fn initialize_result(id: u64) -> serde_json::Value { + json!({ + "jsonrpc": "2.0", + "id": id, + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "fixture", "version": "1.2.3"} + } + }) +} + +/// A single-page `tools/list` result. +fn tools_result(id: u64, tools: serde_json::Value) -> serde_json::Value { + json!({"jsonrpc": "2.0", "id": id, "result": {"tools": tools}}) +} + +fn config(server: &SseTestServer) -> McpSseServerConfig { + McpSseServerConfig::new("fixture", server.sse_url()) +} + +/// Drives the fixture through the handshake so a test can reach a connected +/// client without repeating the three scripted replies. +async fn connect(server: &SseTestServer, config: McpSseServerConfig) -> McpSseClient { + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=abc"); + + server.wait_for_posts(1); + server.send_message(&initialize_result(1)); + + // The initialized notification and tools/list follow. + server.wait_for_posts(3); + server.send_message(&tools_result( + 2, + json!([{ + "name": "search", + "description": "Search the corpus", + "inputSchema": {"type": "object", "properties": {"q": {"type": "string"}}} + }]), + )); + + connecting + .await + .expect("the connect task should not panic") + .expect("the handshake should succeed") +} + +// --------------------------------------------------------------------------- +// Lifecycle +// --------------------------------------------------------------------------- + +#[tokio::test(flavor = "multi_thread")] +async fn completes_the_initialize_initialized_and_tools_list_handshake() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + assert_eq!( + client.server_info().map(|info| info.name.as_str()), + Some("fixture") + ); + assert_eq!(client.tools().len(), 1); + assert_eq!(client.tools()[0].name, "search"); + + let methods: Vec> = server + .posts() + .iter() + .map(|request| request.rpc_method()) + .collect(); + assert_eq!( + methods, + vec![ + Some("initialize".to_string()), + Some("notifications/initialized".to_string()), + Some("tools/list".to_string()), + ], + "the handshake must follow the 2024-11-05 order" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn the_initialized_notification_carries_no_request_id() { + let server = SseTestServer::start(); + let _client = connect(&server, config(&server)).await; + + let notification = server + .posts() + .into_iter() + .find(|request| request.rpc_method().as_deref() == Some("notifications/initialized")) + .expect("the notification should be sent"); + assert!( + notification.rpc_id().is_none(), + "a notification must not carry an id" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn opens_the_stream_with_the_event_stream_accept_header() { + let server = SseTestServer::start(); + let _client = connect(&server, config(&server)).await; + + let stream_request = server + .requests() + .into_iter() + .find(|request| request.method == "GET") + .expect("the stream should be opened with GET"); + assert_eq!( + stream_request.header("accept"), + Some("text/event-stream"), + "the GET must advertise the event stream" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn posts_json_rpc_messages_as_application_json() { + let server = SseTestServer::start(); + let _client = connect(&server, config(&server)).await; + + let post = server.posts().into_iter().next().expect("a POST is sent"); + assert_eq!(post.header("content-type"), Some("application/json")); +} + +#[tokio::test(flavor = "multi_thread")] +async fn posts_to_the_endpoint_named_by_the_server() { + let server = SseTestServer::start(); + let _client = connect(&server, config(&server)).await; + + let post = server.posts().into_iter().next().expect("a POST is sent"); + assert_eq!( + post.target, "/messages/?session_id=abc", + "the session id query must be preserved verbatim" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn accepts_a_202_response_to_the_message_post() { + let server = SseTestServer::start(); + // 202 Accepted is what both reference servers return. + server.queue_post_reply(PostReply::Accepted); + let client = connect(&server, config(&server)).await; + assert_eq!(client.tools().len(), 1); +} + +#[tokio::test(flavor = "multi_thread")] +async fn accepts_a_200_response_to_the_message_post() { + let server = SseTestServer::start(); + server.queue_post_reply(PostReply::Ok); + let client = connect(&server, config(&server)).await; + assert_eq!(client.tools().len(), 1); +} + +#[tokio::test(flavor = "multi_thread")] +async fn walks_every_page_of_a_paginated_tools_list() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=abc"); + server.wait_for_posts(1); + server.send_message(&initialize_result(1)); + + server.wait_for_posts(3); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 2, + "result": { + "tools": [{"name": "first", "inputSchema": {"type": "object"}}], + "nextCursor": "page-2" + } + })); + + server.wait_for_posts(4); + server.send_message(&tools_result( + 3, + json!([{"name": "second", "inputSchema": {"type": "object"}}]), + )); + + let client = connecting + .await + .expect("no panic") + .expect("the handshake should succeed"); + + let names: Vec<&str> = client + .tools() + .iter() + .map(|tool| tool.name.as_str()) + .collect(); + assert_eq!(names, vec!["first", "second"]); + + let cursors: Vec> = server + .posts() + .iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/list")) + .map(|request| { + serde_json::from_str::(&request.body) + .ok()? + .get("params")? + .get("cursor")? + .as_str() + .map(str::to_string) + }) + .collect(); + assert_eq!( + cursors, + vec![None, Some("page-2".to_string())], + "the second page must echo the server's opaque cursor" + ); +} + +/// A server that keeps returning a cursor must not loop forever. Cursors are +/// opaque, so a repeat cannot be detected by value; only a page bound stops it. +#[tokio::test(flavor = "multi_thread")] +async fn stops_paginating_a_server_that_never_ends_its_tools_list() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + max_tool_pages: 4, + ..McpSseLimits::default() + }; + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=abc"); + server.wait_for_posts(1); + server.send_message(&initialize_result(1)); + + // Answer each tools/list with the same cursor. Replies must follow their + // request, not precede it: a response for an unregistered id is dropped. + // With a bound of 4 the client asks exactly four times and then gives up, + // so the count is exact rather than open-ended. + for page in 0..4 { + server.wait_for_posts(3 + page); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 2 + page, + "result": {"tools": [], "nextCursor": "always-more"} + })); + } + + let error = tokio::time::timeout(std::time::Duration::from_secs(10), connecting) + .await + .expect("the client must give up rather than paginate forever") + .expect("no panic") + .expect_err("an endless cursor must fail"); + assert!( + matches!(error, McpSseError::TooManyToolPages { limit: 4 }), + "got {error:?}" + ); + + let pages = server + .posts() + .into_iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/list")) + .count(); + assert_eq!( + pages, 4, + "the client must stop at the configured page bound" + ); +} + +// --------------------------------------------------------------------------- +// Tool calls +// --------------------------------------------------------------------------- + +#[tokio::test(flavor = "multi_thread")] +async fn calls_a_tool_and_returns_its_content() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = tokio::spawn(async move { + ( + client.call_tool("search", Some(json!({"q": "logs"}))).await, + client, + ) + }); + + server.wait_for_posts(4); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": {"content": [{"type": "text", "text": "found it"}], "isError": false} + })); + + let (result, _client) = calling.await.expect("no panic"); + let result = result.expect("the call should succeed"); + assert!(!result.is_error); + assert_eq!(result.content[0].text.as_deref(), Some("found it")); + + let call = server + .posts() + .into_iter() + .find(|request| request.rpc_method().as_deref() == Some("tools/call")) + .expect("the call should be posted"); + let body: serde_json::Value = serde_json::from_str(&call.body).expect("valid JSON"); + assert_eq!(body["params"]["name"], "search"); + assert_eq!(body["params"]["arguments"]["q"], "logs"); +} + +#[tokio::test(flavor = "multi_thread")] +async fn surfaces_a_tool_result_flagged_as_an_error() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = tokio::spawn(async move { (client.call_tool("search", None).await, client) }); + + server.wait_for_posts(4); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": {"content": [{"type": "text", "text": REMOTE_CANARY}], "isError": true} + })); + + let (result, _client) = calling.await.expect("no panic"); + let result = result.expect("an isError result is still a successful response"); + assert!( + result.is_error, + "isError must be preserved rather than turned into a transport failure" + ); + assert_eq!(result.content[0].text.as_deref(), Some(REMOTE_CANARY)); +} + +#[tokio::test(flavor = "multi_thread")] +async fn surfaces_a_json_rpc_error_response() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = tokio::spawn(async move { (client.call_tool("missing", None).await, client) }); + + server.wait_for_posts(4); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "error": { + "code": -32602, + "message": REMOTE_CANARY, + "data": {"forged": REMOTE_CANARY} + } + })); + + let (result, _client) = calling.await.expect("no panic"); + let error = result.expect_err("a JSON-RPC error is a failure"); + let McpSseError::JsonRpc(rpc) = &error else { + panic!("got {error:?}"); + }; + assert_eq!(rpc.code, -32602); + assert_eq!(rpc.message, "server message omitted"); + assert!(rpc.data.is_none(), "server data must be discarded"); + assert_remote_canary_absent(&error); + assert!(error.to_string().contains("-32602"), "got {error}"); +} + +#[tokio::test(flavor = "multi_thread")] +async fn response_decode_errors_do_not_retain_server_text() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = tokio::spawn(async move { (client.call_tool("search", None).await, client) }); + + server.wait_for_posts(4); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": {"content": REMOTE_CANARY, "isError": false} + })); + + let (result, _client) = calling.await.expect("no panic"); + let error = result.expect_err("the response shape is invalid"); + assert!(matches!(error, McpSseError::ParseError(_)), "got {error:?}"); + assert_remote_canary_absent(&error); +} + +#[tokio::test(flavor = "multi_thread")] +async fn resolves_concurrent_calls_whose_responses_arrive_in_reverse_order() { + let server = SseTestServer::start(); + let client = std::sync::Arc::new(connect(&server, config(&server)).await); + + let first = { + let client = std::sync::Arc::clone(&client); + tokio::spawn(async move { client.call_tool("search", Some(json!({"q": "one"}))).await }) + }; + let second = { + let client = std::sync::Arc::clone(&client); + tokio::spawn(async move { client.call_tool("search", Some(json!({"q": "two"}))).await }) + }; + + // Both calls must be in flight before either is answered. + server.wait_for_posts(5); + + // Answer the second request first: the stream carries no ordering guarantee. + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 4, + "result": {"content": [{"type": "text", "text": "second"}], "isError": false} + })); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": {"content": [{"type": "text", "text": "first"}], "isError": false} + })); + + let first = first + .await + .expect("no panic") + .expect("first should resolve"); + let second = second + .await + .expect("no panic") + .expect("second should resolve"); + + assert_eq!( + first.content[0].text.as_deref(), + Some("first"), + "each caller must receive the response matching its own id" + ); + assert_eq!(second.content[0].text.as_deref(), Some("second")); +} + +// --------------------------------------------------------------------------- +// Authentication +// --------------------------------------------------------------------------- + +#[tokio::test(flavor = "multi_thread")] +async fn sends_configured_headers_on_both_the_stream_and_the_posts() { + let server = SseTestServer::start(); + let config = McpSseServerConfig::new("fixture", server.sse_url()) + .with_bearer_token("super-secret-token") + .with_header("x-tenant", "acme") + .allowing_plaintext_credentials(); + let _client = connect(&server, config).await; + + let requests = server.requests(); + let stream_request = requests + .iter() + .find(|request| request.method == "GET") + .expect("the stream is opened"); + assert_eq!( + stream_request.header("authorization"), + Some("Bearer super-secret-token"), + "the GET must carry the credential" + ); + assert_eq!(stream_request.header("x-tenant"), Some("acme")); + + let post = requests + .iter() + .find(|request| request.method == "POST") + .expect("a message is posted"); + assert_eq!( + post.header("authorization"), + Some("Bearer super-secret-token"), + "the POST must carry the credential too" + ); + assert_eq!(post.header("x-tenant"), Some("acme")); +} + +#[tokio::test(flavor = "multi_thread")] +async fn header_values_never_appear_in_client_debug_output() { + let server = SseTestServer::start(); + let config = McpSseServerConfig::new("fixture", server.sse_url()) + .with_bearer_token("super-secret-token") + .allowing_plaintext_credentials(); + let client = connect(&server, config).await; + + let rendered = format!("{client:?}"); + assert!( + !rendered.contains("super-secret-token"), + "the client must not render its credentials: {rendered}" + ); +} + +// --------------------------------------------------------------------------- +// Endpoint handling +// --------------------------------------------------------------------------- + +#[tokio::test(flavor = "multi_thread")] +async fn rejects_an_endpoint_pointing_at_another_origin() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.send_endpoint("https://remote-canary-must-not-surface.invalid/messages"); + + let error = connecting + .await + .expect("no panic") + .expect_err("a cross-origin endpoint must be refused"); + assert!(matches!(error, McpSseError::Endpoint(_)), "got {error:?}"); + assert_remote_canary_absent(&error); + assert!( + !format!("{error:?}").contains("remote-canary-must-not-surface"), + "got {error:?}" + ); + + assert!( + server.posts().is_empty(), + "nothing may be sent once the endpoint is refused" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn rejects_a_protocol_relative_endpoint() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + // Looks like a path but replaces the whole authority. + server.send_endpoint("//evil.example/messages"); + + let error = connecting + .await + .expect("no panic") + .expect_err("a protocol-relative endpoint must be refused"); + assert!(matches!(error, McpSseError::Endpoint(_)), "got {error:?}"); + assert!(server.posts().is_empty()); +} + +#[tokio::test(flavor = "multi_thread")] +async fn honors_only_the_first_endpoint_event() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=first"); + // A later endpoint event must not redirect traffic mid-session. + server.send_endpoint("/messages/?session_id=second"); + + server.wait_for_posts(1); + server.send_message(&initialize_result(1)); + server.wait_for_posts(3); + server.send_message(&tools_result(2, json!([]))); + + let _client = connecting.await.expect("no panic").expect("handshake"); + + for post in server.posts() { + assert_eq!( + post.target, "/messages/?session_id=first", + "every POST must use the first endpoint" + ); + } +} + +#[tokio::test(flavor = "multi_thread")] +async fn rejects_an_oversized_endpoint_event() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + max_endpoint_bytes: 64, + ..McpSseLimits::default() + }; + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.send_endpoint(&format!("/messages/?session_id={}", "x".repeat(512))); + + let error = connecting + .await + .expect("no panic") + .expect_err("an oversized endpoint must be refused"); + assert!( + matches!(error, McpSseError::EndpointTooLarge { limit: 64 }), + "got {error:?}" + ); + assert!(server.posts().is_empty()); +} + +// --------------------------------------------------------------------------- +// Stream framing +// --------------------------------------------------------------------------- + +#[tokio::test(flavor = "multi_thread")] +async fn reassembles_an_endpoint_event_split_across_chunks() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + // One logical event delivered as three separate TCP chunks. + server.send_raw("event: end"); + server.send_raw("point\ndata: /messa"); + server.send_raw("ges/?session_id=abc\n\n"); + + server.wait_for_posts(1); + server.send_message(&initialize_result(1)); + server.wait_for_posts(3); + server.send_message(&tools_result(2, json!([]))); + + let _client = connecting.await.expect("no panic").expect("handshake"); + assert_eq!( + server.posts()[0].target, + "/messages/?session_id=abc", + "a split event must reassemble exactly" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn reads_a_stream_using_crlf_terminators_and_heartbeats() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + // sse-starlette, used by most Python MCP servers, defaults to CRLF and + // sends comment-only heartbeats. + server.send_raw(": ping - keepalive\r\n\r\n"); + server.send_raw("event: endpoint\r\ndata: /messages/?session_id=abc\r\n\r\n"); + + server.wait_for_posts(1); + server.send_raw(": ping - keepalive\r\n\r\n"); + server.send_raw(format!( + "event: message\r\ndata: {}\r\n\r\n", + initialize_result(1) + )); + + server.wait_for_posts(3); + server.send_raw(format!( + "event: message\r\ndata: {}\r\n\r\n", + tools_result( + 2, + json!([{"name": "search", "inputSchema": {"type": "object"}}]) + ) + )); + + let client = connecting.await.expect("no panic").expect("handshake"); + assert_eq!(client.tools().len(), 1); +} + +#[tokio::test(flavor = "multi_thread")] +async fn reads_a_message_split_across_several_data_lines() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = tokio::spawn(async move { (client.call_tool("search", None).await, client) }); + + server.wait_for_posts(4); + // A JSON payload containing newlines is emitted as multiple data lines. + server.send_raw( + "event: message\ndata: {\"jsonrpc\":\"2.0\",\"id\":3,\"result\":\ndata: {\"content\":[{\"type\":\"text\",\"text\":\"ok\"}],\"isError\":false}}\n\n", + ); + + let (result, _client) = calling.await.expect("no panic"); + let result = result.expect("multi-line data must rejoin into one payload"); + assert_eq!(result.content[0].text.as_deref(), Some("ok")); +} + +#[tokio::test(flavor = "multi_thread")] +async fn ignores_unknown_event_names_such_as_ping() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + // Older sse-starlette emits a real event with non-JSON data. + server.send_raw("event: ping\ndata: 2026-08-08 12:00:00\n\n"); + server.send_endpoint("/messages/?session_id=abc"); + + server.wait_for_posts(1); + server.send_raw("event: ping\ndata: 2026-08-08 12:00:15\n\n"); + server.send_message(&initialize_result(1)); + server.wait_for_posts(3); + server.send_message(&tools_result(2, json!([]))); + + let client = connecting.await.expect("no panic").expect("handshake"); + assert!(client.tools().is_empty()); +} + +#[tokio::test(flavor = "multi_thread")] +async fn ignores_a_server_initiated_request_rather_than_treating_it_as_a_response() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = tokio::spawn(async move { (client.call_tool("search", None).await, client) }); + + server.wait_for_posts(4); + // A ping request carries method and id but is not a response to id 3. + server.send_message(&json!({"jsonrpc": "2.0", "id": 3, "method": "ping"})); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": {"content": [{"type": "text", "text": "real"}], "isError": false} + })); + + let (result, _client) = calling.await.expect("no panic"); + let result = result.expect("the real response must still resolve the call"); + assert_eq!(result.content[0].text.as_deref(), Some("real")); +} + +#[tokio::test(flavor = "multi_thread")] +async fn ignores_a_repeated_response_for_an_already_answered_id() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = tokio::spawn(async move { (client.call_tool("search", None).await, client) }); + + server.wait_for_posts(4); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": {"content": [{"type": "text", "text": "first"}], "isError": false} + })); + // A second result for the same id must not reach the caller. + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": {"content": [{"type": "text", "text": "second"}], "isError": false} + })); + + let (result, client) = calling.await.expect("no panic"); + assert_eq!( + result.expect("the first response wins").content[0] + .text + .as_deref(), + Some("first") + ); + + // The connection stays usable rather than being corrupted by the duplicate. + client.shutdown().await; +} + +#[tokio::test(flavor = "multi_thread")] +async fn ignores_malformed_json_rpc_without_failing_other_calls() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = tokio::spawn(async move { (client.call_tool("search", None).await, client) }); + + server.wait_for_posts(4); + server.send_raw("event: message\ndata: {not json at all\n\n"); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 3, + "result": {"content": [{"type": "text", "text": "ok"}], "isError": false} + })); + + let (result, _client) = calling.await.expect("no panic"); + assert_eq!( + result + .expect("a malformed frame must not break the stream") + .content[0] + .text + .as_deref(), + Some("ok") + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn tears_down_the_stream_when_an_event_exceeds_the_size_limit() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + max_event_bytes: 256, + ..McpSseLimits::default() + }; + let client = connect(&server, config).await; + + let calling = tokio::spawn(async move { (client.call_tool("search", None).await, client) }); + + server.wait_for_posts(4); + server.send_raw(format!("event: message\ndata: {}\n\n", "x".repeat(4096))); + + let (result, _client) = calling.await.expect("no panic"); + let error = result.expect_err("an oversized event must fail the call"); + assert!( + matches!(error, McpSseError::RequestIndeterminate { .. }), + "an accepted call that never answered is indeterminate, got {error:?}" + ); +} + +// --------------------------------------------------------------------------- +// Connection failures +// --------------------------------------------------------------------------- + +#[tokio::test(flavor = "multi_thread")] +async fn rejects_a_stream_response_that_is_not_an_event_stream() { + let server = SseTestServer::with_opening(StreamOpening::WrongContentType); + let error = McpSseClient::connect(&config(&server)) + .await + .expect_err("a JSON response is not a stream"); + assert!( + matches!(error, McpSseError::UnexpectedContentType { .. }), + "got {error:?}" + ); + assert!(!error.to_string().contains("remote-canary"), "got {error}"); + assert!( + !format!("{error:?}").contains("remote-canary"), + "got {error:?}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn accepts_an_event_stream_content_type_carrying_a_charset() { + // The fixture answers `text/event-stream; charset=utf-8`, which is what + // real servers send; a strict equality check would reject it. + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + assert_eq!(client.tools().len(), 1); +} + +#[tokio::test(flavor = "multi_thread")] +async fn rejects_a_non_success_status_on_the_stream() { + let server = SseTestServer::with_opening(StreamOpening::Status { + code: 404, + body: "not found".to_string(), + }); + let error = McpSseClient::connect(&config(&server)) + .await + .expect_err("404 is not a stream"); + assert!( + matches!(error, McpSseError::HttpStatus { status, .. } if status == 404), + "got {error:?}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn refuses_to_follow_a_redirect_on_the_stream() { + let server = SseTestServer::with_opening(StreamOpening::Redirect { + location: "http://evil.example/sse".to_string(), + }); + let error = McpSseClient::connect(&config(&server)) + .await + .expect_err("a redirect must not be followed"); + assert!( + matches!(error, McpSseError::RedirectRefused), + "got {error:?}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn reports_a_rejected_post_without_quoting_the_response_body() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.queue_post_reply(PostReply::Status { + code: 400, + body: "SESSION-SECRET-LEAK".to_string(), + }); + server.send_endpoint("/messages/?session_id=abc"); + + let error = connecting + .await + .expect("no panic") + .expect_err("a 400 fails the handshake"); + assert!( + matches!(error, McpSseError::HttpStatus { status, .. } if status == 400), + "got {error:?}" + ); + assert!( + !error.to_string().contains("SESSION-SECRET-LEAK"), + "server text must never reach an error: {error}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn reports_a_server_error_on_the_message_post() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.queue_post_reply(PostReply::Status { + code: 503, + body: "unavailable".to_string(), + }); + server.send_endpoint("/messages/?session_id=abc"); + + let error = connecting + .await + .expect("no panic") + .expect_err("a 503 fails the handshake"); + assert!( + matches!(error, McpSseError::HttpStatus { status, .. } if status == 503), + "got {error:?}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn refuses_to_follow_a_redirect_on_a_message_post() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.queue_post_reply(PostReply::Redirect { + location: "http://evil.example/messages".to_string(), + }); + server.send_endpoint("/messages/?session_id=abc"); + + let error = connecting + .await + .expect("no panic") + .expect_err("a redirected POST must not be followed"); + assert!( + matches!(error, McpSseError::RedirectRefused), + "got {error:?}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn rejects_an_absolute_endpoint_on_the_fixture_origin_with_a_different_port() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + // Same host, different port: still a different origin. + let other_port = server + .base_url() + .rsplit(':') + .next() + .and_then(|port| port.parse::().ok()) + .map(|port| port.wrapping_add(1)) + .expect("the fixture URL carries a port"); + server.send_endpoint(&format!("http://127.0.0.1:{other_port}/messages/")); + + let error = connecting + .await + .expect("no panic") + .expect_err("a different port is a different origin"); + assert!(matches!(error, McpSseError::Endpoint(_)), "got {error:?}"); + assert!(server.posts().is_empty()); +} + +#[tokio::test(flavor = "multi_thread")] +async fn accepts_an_absolute_endpoint_on_the_configured_origin() { + let server = SseTestServer::start(); + let base_url = server.base_url().to_string(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + // The specification permits an absolute URL as long as the origin matches. + server.send_endpoint(&format!("{base_url}/messages/?session_id=abc")); + + server.wait_for_posts(1); + server.send_message(&initialize_result(1)); + server.wait_for_posts(3); + server.send_message(&tools_result(2, json!([]))); + + let _client = connecting.await.expect("no panic").expect("handshake"); + assert_eq!(server.posts()[0].target, "/messages/?session_id=abc"); +} + +#[tokio::test(flavor = "multi_thread")] +async fn reports_a_post_the_server_never_answers() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + // The server accepts the connection then closes it without responding. + server.queue_post_reply(PostReply::Drop); + server.send_endpoint("/messages/?session_id=abc"); + + let error = connecting + .await + .expect("no panic") + .expect_err("a dropped POST fails the handshake"); + assert!( + matches!(error, McpSseError::Transport(_)), + "a dropped connection is a transport failure, got {error:?}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn bounds_the_initialize_request_post() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + initialize_timeout: std::time::Duration::from_millis(150), + ..McpSseLimits::default() + }; + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.queue_post_reply(PostReply::StallBeforeHeaders); + server.send_endpoint("/messages/?session_id=abc"); + server.wait_for_posts(1); + + let error = tokio::time::timeout(std::time::Duration::from_secs(2), connecting) + .await + .expect("the configured initialize deadline must include its POST response head") + .expect("no panic") + .expect_err("the initialize POST never receives response headers"); + server.release_stalled_posts(); + + assert!(matches!(error, McpSseError::Timeout(_)), "got {error:?}"); + assert_eq!( + server.posts().len(), + 1, + "an initialize timeout must not send a second request" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn bounds_the_initialized_notification_post() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + initialize_timeout: std::time::Duration::from_millis(150), + ..McpSseLimits::default() + }; + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.send_endpoint("/messages/?session_id=abc"); + server.wait_for_posts(1); + server.queue_post_reply(PostReply::StallBeforeHeaders); + server.send_message(&initialize_result(1)); + server.wait_for_posts(2); + + let error = tokio::time::timeout(std::time::Duration::from_secs(2), connecting) + .await + .expect("the configured initialize deadline must bound the notification POST") + .expect("no panic") + .expect_err("a notification POST that never answers must fail connect"); + server.release_stalled_posts(); + + assert!(matches!(error, McpSseError::Timeout(_)), "got {error:?}"); + assert_eq!( + server.posts().len(), + 2, + "tools/list must not start after the initialized notification timed out" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn bounds_the_entire_tool_call_when_post_headers_never_arrive() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + call_tool_timeout: std::time::Duration::from_millis(150), + ..McpSseLimits::default() + }; + let client = std::sync::Arc::new(connect(&server, config).await); + + server.queue_post_reply(PostReply::StallBeforeHeaders); + let calling = { + let client = std::sync::Arc::clone(&client); + tokio::spawn(async move { client.call_tool("charge_card", None).await }) + }; + server.wait_for_posts(4); + + let error = tokio::time::timeout(std::time::Duration::from_secs(2), calling) + .await + .expect("the configured call deadline must include the POST response head") + .expect("no panic") + .expect_err("the server withheld its response head"); + assert!( + matches!(error, McpSseError::RequestIndeterminate { .. }), + "the server read the request body, so delivery is ambiguous: {error:?}" + ); + assert_eq!( + server + .posts() + .iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/call")) + .count(), + 1, + "the ambiguous tool call must never be replayed" + ); + + // The timeout removes only this request's correlation state. Once the + // fixture releases the abandoned POST connection, the SSE session remains + // able to correlate a later, explicitly requested call. + server.release_stalled_posts(); + let calling = { + let client = std::sync::Arc::clone(&client); + tokio::spawn(async move { client.call_tool("search", None).await }) + }; + server.wait_for_posts(5); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 4, + "result": {"content": [{"type": "text", "text": "still usable"}], "isError": false} + })); + assert_eq!( + calling + .await + .expect("no panic") + .expect("a later explicit call should succeed") + .content[0] + .text + .as_deref(), + Some("still usable") + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn bounds_the_entire_tool_call_while_draining_the_post_body() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + call_tool_timeout: std::time::Duration::from_millis(150), + ..McpSseLimits::default() + }; + let client = connect(&server, config).await; + + server.queue_post_reply(PostReply::StallAfterHeaders); + let calling = + tokio::spawn(async move { (client.call_tool("charge_card", None).await, client) }); + server.wait_for_posts(4); + server.wait_for_post_response_headers(4); + + let (result, _client) = tokio::time::timeout(std::time::Duration::from_secs(2), calling) + .await + .expect("the configured call deadline must include response-body drain") + .expect("no panic"); + server.release_stalled_posts(); + + let error = result.expect_err("the declared response body never arrived"); + assert!( + matches!(error, McpSseError::RequestIndeterminate { .. }), + "the tool may have run before the POST response body stalled: {error:?}" + ); + assert_eq!( + server + .posts() + .iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/call")) + .count(), + 1, + "the ambiguous tool call must never be replayed" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn a_tool_call_dropped_after_its_body_is_read_is_indeterminate() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + server.queue_post_reply(PostReply::Drop); + let error = client + .call_tool("charge_card", None) + .await + .expect_err("the fixture drops the POST connection after reading its body"); + + assert!( + matches!(error, McpSseError::RequestIndeterminate { .. }), + "a transport failure cannot prove non-delivery: {error:?}" + ); + assert_eq!( + server + .posts() + .iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/call")) + .count(), + 1, + "the dropped tool call must never be replayed" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn a_tool_call_answered_with_an_http_error_is_indeterminate() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + server.queue_post_reply(PostReply::Status { + code: 503, + body: "failed after dispatch".to_string(), + }); + let error = client + .call_tool("charge_card", None) + .await + .expect_err("a non-success POST status fails the call"); + + assert!( + matches!(error, McpSseError::RequestIndeterminate { .. }), + "HTTP status cannot prove that the server did no work: {error:?}" + ); + assert_eq!( + server + .posts() + .iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/call")) + .count(), + 1, + "the failed tool call must never be replayed" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn a_tool_call_answered_with_a_redirect_is_indeterminate_and_not_followed() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + server.queue_post_reply(PostReply::Redirect { + location: "http://evil.example/messages".to_string(), + }); + let error = client + .call_tool("charge_card", None) + .await + .expect_err("a redirected POST fails the call"); + + assert!( + matches!(error, McpSseError::RequestIndeterminate { .. }), + "a redirect response cannot prove that the original server did no work: {error:?}" + ); + assert_eq!( + server + .posts() + .iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/call")) + .count(), + 1, + "the client must neither follow nor replay the redirected tool call" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn cancelling_a_tool_call_future_removes_its_pending_waiter() { + let server = SseTestServer::start(); + let client = std::sync::Arc::new(connect(&server, config(&server)).await); + + server.queue_post_reply(PostReply::StallBeforeHeaders); + let calling = { + let client = std::sync::Arc::clone(&client); + tokio::spawn(async move { client.call_tool("charge_card", None).await }) + }; + server.wait_for_posts(4); + calling.abort(); + assert!( + calling + .await + .expect_err("the call task was cancelled") + .is_cancelled() + ); + + assert!( + super::lock_pending(&client.pending).waiters.is_empty(), + "cancelling the future must remove its pending correlation entry" + ); + assert_eq!( + server + .posts() + .iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/call")) + .count(), + 1, + "cancellation must not cause an automatic replay" + ); + server.release_stalled_posts(); +} + +// --------------------------------------------------------------------------- +// Teardown +// --------------------------------------------------------------------------- + +#[tokio::test(flavor = "multi_thread")] +async fn fails_every_pending_call_when_the_stream_reaches_eof() { + let server = SseTestServer::start(); + let client = std::sync::Arc::new(connect(&server, config(&server)).await); + + let first = { + let client = std::sync::Arc::clone(&client); + tokio::spawn(async move { client.call_tool("search", Some(json!({"q": "a"}))).await }) + }; + let second = { + let client = std::sync::Arc::clone(&client); + tokio::spawn(async move { client.call_tool("search", Some(json!({"q": "b"}))).await }) + }; + + server.wait_for_posts(5); + server.close_stream(); + + let first = first + .await + .expect("no panic") + .expect_err("EOF must fail the call rather than hang"); + let second = second + .await + .expect("no panic") + .expect_err("EOF must fail every call"); + + assert!( + matches!(first, McpSseError::RequestIndeterminate { .. }), + "got {first:?}" + ); + assert!( + matches!(second, McpSseError::RequestIndeterminate { .. }), + "got {second:?}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn an_accepted_call_lost_to_a_stream_drop_is_reported_as_indeterminate() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = + tokio::spawn(async move { (client.call_tool("charge_card", None).await, client) }); + + server.wait_for_posts(4); + // The POST was accepted, so the tool may well have run. + server.abort_stream(); + + let (result, _client) = calling.await.expect("no panic"); + let error = result.expect_err("a lost response is a failure"); + match error { + McpSseError::RequestIndeterminate { method } => assert_eq!(method, "tools/call"), + other => panic!("an accepted-but-unanswered call must be indeterminate, got {other:?}"), + } + assert!( + error_says_do_not_retry(&McpSseError::RequestIndeterminate { + method: "tools/call".to_string() + }), + "the message must warn against automatic retry" + ); +} + +fn error_says_do_not_retry(error: &McpSseError) -> bool { + let rendered = error.to_string(); + rendered.contains("may have executed") && rendered.contains("must not be retried") +} + +#[tokio::test(flavor = "multi_thread")] +async fn never_replays_a_tool_call_after_an_ambiguous_failure() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + + let calling = + tokio::spawn(async move { (client.call_tool("charge_card", None).await, client) }); + + server.wait_for_posts(4); + server.abort_stream(); + + let (result, _client) = calling.await.expect("no panic"); + result.expect_err("the call fails"); + + // Give any (incorrect) retry a chance to appear before asserting. + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + + let calls = server + .posts() + .into_iter() + .filter(|request| request.rpc_method().as_deref() == Some("tools/call")) + .count(); + assert_eq!( + calls, 1, + "a tools/call may have side effects and must never be re-sent" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn shutting_down_fails_calls_that_are_still_in_flight() { + let server = SseTestServer::start(); + let client = std::sync::Arc::new(connect(&server, config(&server)).await); + + let calling = { + let client = std::sync::Arc::clone(&client); + tokio::spawn(async move { client.call_tool("search", None).await }) + }; + + server.wait_for_posts(4); + client.shutdown().await; + + let error = calling + .await + .expect("no panic") + .expect_err("shutdown must resolve outstanding calls"); + assert!( + matches!(error, McpSseError::RequestIndeterminate { .. }), + "got {error:?}" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn a_request_made_after_shutdown_fails_immediately() { + let server = SseTestServer::start(); + let client = connect(&server, config(&server)).await; + client.shutdown().await; + + let error = client + .call_tool("search", None) + .await + .expect_err("a shut-down client accepts no work"); + assert!(matches!(error, McpSseError::StreamClosed), "got {error:?}"); +} + +#[tokio::test(flavor = "multi_thread")] +async fn a_timed_out_request_does_not_leak_its_pending_entry() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + call_tool_timeout: std::time::Duration::from_millis(150), + ..McpSseLimits::default() + }; + let client = connect(&server, config).await; + + // Time out several calls, then confirm a later one still succeeds. + for _ in 0..3 { + let error = client + .call_tool("search", None) + .await + .expect_err("no response arrives"); + assert!( + matches!(error, McpSseError::RequestIndeterminate { .. }), + "got {error:?}" + ); + } + + let calling = tokio::spawn(async move { (client.call_tool("search", None).await, client) }); + server.wait_for_posts(7); + server.send_message(&json!({ + "jsonrpc": "2.0", + "id": 6, + "result": {"content": [{"type": "text", "text": "late but fine"}], "isError": false} + })); + + let (result, _client) = calling.await.expect("no panic"); + assert_eq!( + result.expect("the connection remains usable").content[0] + .text + .as_deref(), + Some("late but fine") + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn connecting_times_out_when_the_endpoint_event_never_arrives() { + let server = SseTestServer::start(); + let mut config = config(&server); + config.limits = McpSseLimits { + connect_timeout: std::time::Duration::from_millis(200), + ..McpSseLimits::default() + }; + + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + server.wait_for_stream(); + // A buffering proxy is the common cause; no endpoint event is ever sent. + + let error = connecting + .await + .expect("no panic") + .expect_err("the handshake cannot proceed without an endpoint"); + assert!(matches!(error, McpSseError::Timeout(_)), "got {error:?}"); +} + +#[tokio::test(flavor = "multi_thread")] +async fn connecting_fails_when_the_stream_closes_before_the_endpoint_arrives() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.close_stream(); + + let error = connecting + .await + .expect("no panic") + .expect_err("a closed stream cannot complete the handshake"); + assert!(matches!(error, McpSseError::StreamClosed), "got {error:?}"); +} + +/// A rejected endpoint must not leave the reader task running. +/// +/// The task owns the response body, so leaking it also leaks the connection. +/// Because the reader is what consumes the stream, a leaked one keeps draining +/// events the abandoned client can never deliver — observable here as the +/// fixture continuing to accept writes long after connect returned. +#[tokio::test(flavor = "multi_thread")] +async fn a_refused_endpoint_leaves_no_reader_consuming_the_stream() { + let server = SseTestServer::start(); + let config = config(&server); + let connecting = tokio::spawn(async move { McpSseClient::connect(&config).await }); + + server.wait_for_stream(); + server.send_endpoint("https://evil.example/messages"); + + let error = connecting + .await + .expect("no panic") + .expect_err("a cross-origin endpoint is refused"); + assert!(matches!(error, McpSseError::Endpoint(_)), "got {error:?}"); + + // Nothing was ever sent to the server, and nothing may be sent later. + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + assert!( + server.posts().is_empty(), + "a refused endpoint must not produce any request" + ); +} diff --git a/vendor/mentra/src/mcp/sse/config.rs b/vendor/mentra/src/mcp/sse/config.rs new file mode 100644 index 0000000..2b70c65 --- /dev/null +++ b/vendor/mentra/src/mcp/sse/config.rs @@ -0,0 +1,280 @@ +//! Configuration for the legacy MCP HTTP+SSE transport. + +#[cfg(test)] +mod tests; + +use std::collections::BTreeMap; +use std::time::Duration; + +use serde::Deserialize; +use url::Url; + +use super::endpoint::{EndpointError, validate_stream_url}; + +/// Default timeout for opening the SSE stream and reading its response head. +pub const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +/// Default timeout for the MCP `initialize` handshake, matching the stdio client. +pub const DEFAULT_INITIALIZE_TIMEOUT: Duration = Duration::from_secs(10); +/// Default timeout for `tools/list`, matching the stdio client. +pub const DEFAULT_LIST_TOOLS_TIMEOUT: Duration = Duration::from_secs(30); +/// Default timeout for `tools/call`, matching the stdio client. +pub const DEFAULT_CALL_TOOL_TIMEOUT: Duration = Duration::from_secs(120); +/// Default idle timeout between stream reads. +/// +/// Servers built on `sse-starlette` — which covers most Python MCP servers — +/// emit a comment heartbeat every 15 seconds, so five minutes of silence means +/// the stream is dead rather than quiet. +pub const DEFAULT_STREAM_IDLE_TIMEOUT: Duration = Duration::from_secs(300); +/// Default cap on the bytes buffered for a single SSE event. +/// +/// The largest legitimate event is a `tools/call` result carrying base64 +/// content; 4 MiB of base64 is roughly 3 MB of binary, which is generous for a +/// tool result while bounding worst-case memory per connection. +pub const DEFAULT_MAX_EVENT_BYTES: usize = 4 * 1024 * 1024; +/// Default cap on the bytes buffered for the `endpoint` event specifically. +/// +/// The endpoint event is processed before any request correlation exists, so it +/// is the earliest attacker-reachable allocation in the connection lifecycle. +/// Its payload is a single URL, and common proxy header limits sit near 2 KiB. +pub const DEFAULT_MAX_ENDPOINT_BYTES: usize = 8 * 1024; +/// Default cap on how many `tools/list` pages are followed. +/// +/// Cursors are opaque, so a server repeating one cannot be detected by value; +/// only a page bound stops the walk. A server needing more pages than this to +/// describe its tools is malfunctioning. +pub const DEFAULT_MAX_TOOL_PAGES: usize = 1_000; + +/// A header value that is never rendered by `Debug` or `Display`. +/// +/// Redaction is a property of this type rather than of each container, so every +/// struct that derives `Debug` inherits it without a rule for contributors to +/// remember. +/// +/// This type deliberately does **not** implement [`serde::Serialize`]. Adding +/// `#[derive(Serialize)]` to any struct holding one is therefore a compile +/// error rather than a silent credential leak into a config dump, a state +/// snapshot, or a session-persistence layer. +#[derive(Clone, PartialEq, Eq, Deserialize)] +#[serde(transparent)] +pub struct SecretString(String); + +impl SecretString { + /// Wraps a value that must not be logged. + pub fn new(value: impl Into) -> Self { + Self(value.into()) + } + + /// Returns the wrapped value. + /// + /// This is the single grep-able point at which a secret becomes visible. + pub fn expose_secret(&self) -> &str { + &self.0 + } +} + +impl std::fmt::Debug for SecretString { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("SecretString([redacted])") + } +} + +impl> From for SecretString { + fn from(value: T) -> Self { + Self::new(value) + } +} + +/// Timeouts and size limits for one SSE connection. +/// +/// These are separated from the operator-facing fields of +/// [`McpSseServerConfig`] because they are tuning knobs for the host rather +/// than something an operator writes in a configuration file. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct McpSseLimits { + /// Bound on opening the stream and reading its response head. + pub connect_timeout: Duration, + /// Bound on the `initialize` handshake. + pub initialize_timeout: Duration, + /// Bound on each `tools/list` page. + pub list_tools_timeout: Duration, + /// Bound on each `tools/call`. + pub call_tool_timeout: Duration, + /// Bound on silence between reads on the SSE stream. + pub stream_idle_timeout: Duration, + /// Bound on the bytes buffered for a single SSE event. + pub max_event_bytes: usize, + /// Bound on the bytes buffered for the `endpoint` event. + pub max_endpoint_bytes: usize, + /// Bound on how many `tools/list` pages are followed. + pub max_tool_pages: usize, +} + +impl Default for McpSseLimits { + fn default() -> Self { + Self { + connect_timeout: DEFAULT_CONNECT_TIMEOUT, + initialize_timeout: DEFAULT_INITIALIZE_TIMEOUT, + list_tools_timeout: DEFAULT_LIST_TOOLS_TIMEOUT, + call_tool_timeout: DEFAULT_CALL_TOOL_TIMEOUT, + stream_idle_timeout: DEFAULT_STREAM_IDLE_TIMEOUT, + max_event_bytes: DEFAULT_MAX_EVENT_BYTES, + max_endpoint_bytes: DEFAULT_MAX_ENDPOINT_BYTES, + max_tool_pages: DEFAULT_MAX_TOOL_PAGES, + } + } +} + +/// Configuration for an MCP server reachable over the legacy HTTP+SSE +/// transport. +/// +/// This is the SSE counterpart to [`McpServerConfig`](crate::mcp::McpServerConfig), +/// which remains the stdio configuration type. +/// +/// # Example +/// +/// ```rust +/// use mentra::mcp::McpSseServerConfig; +/// +/// let config = McpSseServerConfig::new("observability", "https://mcp.example.com/sse") +/// .with_header("authorization", "Bearer "); +/// ``` +/// +/// # Security +/// +/// Header values are stored as [`SecretString`] and never appear in `Debug` +/// output, error messages, or logs. Configuring a header on a plaintext `http://` +/// URL is rejected unless the host is loopback, because the credential would +/// otherwise cross the network in the clear; see +/// [`allow_plaintext_credentials`](Self::allow_plaintext_credentials) to +/// override that deliberately. +#[derive(Debug, Clone, Deserialize)] +pub struct McpSseServerConfig { + /// Display name for the server, used to namespace its bridged tools. + pub name: String, + /// The operator-configured SSE stream URL, opened with a long-lived `GET`. + pub url: String, + /// Headers sent on both the SSE `GET` and every JSON-RPC `POST`. + #[serde(default)] + pub headers: BTreeMap, + /// Permits sending configured headers over plaintext `http://` to a + /// non-loopback host. + /// + /// This exists so the refusal is overridable but never accidental. + #[serde(default)] + pub allow_plaintext_credentials: bool, + /// Timeouts and size limits, defaulted rather than deserialized. + #[serde(skip)] + pub limits: McpSseLimits, +} + +/// Errors from validating an [`McpSseServerConfig`]. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum McpSseConfigError { + #[error("invalid MCP SSE stream URL: {0}")] + Url(#[from] EndpointError), + + #[error("MCP SSE server name must not be empty")] + EmptyName, + + #[error("invalid MCP SSE header name '{name}'")] + InvalidHeaderName { name: String }, + + /// Rendered without the value so a malformed credential never reaches a log. + #[error("MCP SSE header '{name}' has a value that is not valid for HTTP")] + InvalidHeaderValue { name: String }, + + #[error( + "refusing to send configured headers to '{url}' over plaintext http; \ + use https, a loopback host, or set allow_plaintext_credentials" + )] + PlaintextCredentials { url: String }, +} + +impl McpSseServerConfig { + /// Creates a configuration with default timeouts and limits. + pub fn new(name: impl Into, url: impl Into) -> Self { + Self { + name: name.into(), + url: url.into(), + headers: BTreeMap::new(), + allow_plaintext_credentials: false, + limits: McpSseLimits::default(), + } + } + + /// Adds a header sent on both the SSE `GET` and every JSON-RPC `POST`. + pub fn with_header(mut self, name: impl Into, value: impl Into) -> Self { + self.headers.insert(name.into(), value.into()); + self + } + + /// Adds a bearer `Authorization` header. + pub fn with_bearer_token(self, token: impl Into) -> Self { + self.with_header( + "authorization", + SecretString::new(format!("Bearer {}", token.into())), + ) + } + + /// Replaces the timeouts and size limits. + pub fn with_limits(mut self, limits: McpSseLimits) -> Self { + self.limits = limits; + self + } + + /// Permits sending configured headers over plaintext `http://`. + pub fn allowing_plaintext_credentials(mut self) -> Self { + self.allow_plaintext_credentials = true; + self + } + + /// Validates the configuration and returns the parsed stream URL. + /// + /// Runs before any connection is opened so a bad configuration fails at the + /// boundary rather than mid-handshake. + /// Checks the URL, header names, and credential handling without + /// connecting, so a host can reject a bad configuration at its own + /// boundary rather than discovering it mid-build. + pub fn validate(&self) -> Result { + if self.name.trim().is_empty() { + return Err(McpSseConfigError::EmptyName); + } + + let url = validate_stream_url(&self.url)?; + + for (name, value) in &self.headers { + if reqwest::header::HeaderName::try_from(name.as_str()).is_err() { + return Err(McpSseConfigError::InvalidHeaderName { + name: name.to_string(), + }); + } + if reqwest::header::HeaderValue::try_from(value.expose_secret()).is_err() { + return Err(McpSseConfigError::InvalidHeaderValue { + name: name.to_string(), + }); + } + } + + if !self.headers.is_empty() + && url.scheme() == "http" + && !self.allow_plaintext_credentials + && !is_loopback(&url) + { + return Err(McpSseConfigError::PlaintextCredentials { + url: self.url.clone(), + }); + } + + Ok(url) + } +} + +/// Reports whether a URL addresses the loopback interface. +fn is_loopback(url: &Url) -> bool { + match url.host() { + Some(url::Host::Domain(host)) => host.eq_ignore_ascii_case("localhost"), + Some(url::Host::Ipv4(address)) => address.is_loopback(), + Some(url::Host::Ipv6(address)) => address.is_loopback(), + None => false, + } +} diff --git a/vendor/mentra/src/mcp/sse/config/tests.rs b/vendor/mentra/src/mcp/sse/config/tests.rs new file mode 100644 index 0000000..279d82c --- /dev/null +++ b/vendor/mentra/src/mcp/sse/config/tests.rs @@ -0,0 +1,243 @@ +//! Tests for SSE server configuration and secret redaction. + +use super::{McpSseConfigError, McpSseServerConfig, SecretString}; + +// --------------------------------------------------------------------------- +// Secret redaction +// --------------------------------------------------------------------------- + +#[test] +fn secret_debug_output_hides_the_value() { + let secret = SecretString::new("Bearer super-secret-token"); + let rendered = format!("{secret:?}"); + assert!(!rendered.contains("super-secret-token"), "got {rendered}"); + assert_eq!(rendered, "SecretString([redacted])"); +} + +#[test] +fn secret_alternate_debug_output_hides_the_value() { + let secret = SecretString::new("Bearer super-secret-token"); + let rendered = format!("{secret:#?}"); + assert!(!rendered.contains("super-secret-token"), "got {rendered}"); +} + +#[test] +fn exposing_a_secret_returns_the_original_value() { + let secret = SecretString::new("Bearer token"); + assert_eq!(secret.expose_secret(), "Bearer token"); +} + +#[test] +fn config_debug_redacts_header_values_but_keeps_names() { + let config = McpSseServerConfig::new("obs", "https://mcp.example.com/sse") + .with_header("authorization", "Bearer super-secret-token") + .with_header("x-tenant", "acme"); + let rendered = format!("{config:?}"); + + assert!( + !rendered.contains("super-secret-token"), + "the token must not appear: {rendered}" + ); + assert!( + !rendered.contains("acme"), + "no header value may appear: {rendered}" + ); + assert!( + rendered.contains("authorization"), + "header names stay visible for diagnosis: {rendered}" + ); + assert!(rendered.contains("x-tenant"), "got {rendered}"); + assert!(rendered.contains("mcp.example.com"), "got {rendered}"); +} + +#[test] +fn config_alternate_debug_redacts_header_values() { + let config = McpSseServerConfig::new("obs", "https://mcp.example.com/sse") + .with_bearer_token("super-secret-token"); + let rendered = format!("{config:#?}"); + assert!(!rendered.contains("super-secret-token"), "got {rendered}"); +} + +#[test] +fn a_bearer_token_is_stored_as_an_authorization_header() { + let config = + McpSseServerConfig::new("obs", "https://mcp.example.com/sse").with_bearer_token("abc123"); + let value = config + .headers + .get("authorization") + .expect("bearer token sets the authorization header"); + assert_eq!(value.expose_secret(), "Bearer abc123"); +} + +// --------------------------------------------------------------------------- +// Validation +// --------------------------------------------------------------------------- + +#[test] +fn accepts_an_https_url_with_headers() { + let config = + McpSseServerConfig::new("obs", "https://mcp.example.com/sse").with_bearer_token("abc123"); + let url = config.validate().expect("https with headers is allowed"); + assert_eq!(url.host_str(), Some("mcp.example.com")); +} + +#[test] +fn accepts_a_plain_http_url_without_headers() { + let config = McpSseServerConfig::new("local", "http://internal.corp:8080/sse"); + config + .validate() + .expect("plaintext without credentials is allowed"); +} + +#[test] +fn rejects_plaintext_credentials_to_a_remote_host() { + let config = + McpSseServerConfig::new("obs", "http://internal.corp/sse").with_bearer_token("abc123"); + let error = config + .validate() + .expect_err("a token must not cross the network in the clear"); + assert!(matches!( + error, + McpSseConfigError::PlaintextCredentials { .. } + )); +} + +#[test] +fn allows_plaintext_credentials_to_localhost() { + let config = + McpSseServerConfig::new("local", "http://localhost:3000/sse").with_bearer_token("abc123"); + config + .validate() + .expect("loopback never leaves the machine"); +} + +#[test] +fn allows_plaintext_credentials_to_the_ipv4_loopback_address() { + let config = + McpSseServerConfig::new("local", "http://127.0.0.1:3000/sse").with_bearer_token("abc"); + config.validate().expect("127.0.0.1 is loopback"); +} + +#[test] +fn allows_plaintext_credentials_to_the_ipv6_loopback_address() { + let config = McpSseServerConfig::new("local", "http://[::1]:3000/sse").with_bearer_token("abc"); + config.validate().expect("::1 is loopback"); +} + +#[test] +fn allows_plaintext_credentials_when_explicitly_opted_in() { + let config = McpSseServerConfig::new("obs", "http://internal.corp/sse") + .with_bearer_token("abc123") + .allowing_plaintext_credentials(); + config + .validate() + .expect("the operator may override deliberately"); +} + +#[test] +fn rejects_an_empty_server_name() { + let config = McpSseServerConfig::new(" ", "https://mcp.example.com/sse"); + let error = config.validate().expect_err("a name is required"); + assert!(matches!(error, McpSseConfigError::EmptyName)); +} + +#[test] +fn rejects_an_unsupported_url_scheme() { + let config = McpSseServerConfig::new("obs", "ws://mcp.example.com/sse"); + let error = config + .validate() + .expect_err("only http and https are allowed"); + assert!(matches!(error, McpSseConfigError::Url(_))); +} + +#[test] +fn rejects_a_url_with_embedded_credentials() { + let config = McpSseServerConfig::new("obs", "https://user:pass@mcp.example.com/sse"); + let error = config + .validate() + .expect_err("credentials belong in headers, not the URL"); + assert!(matches!(error, McpSseConfigError::Url(_))); +} + +#[test] +fn rejects_a_header_name_that_is_not_valid_for_http() { + let config = McpSseServerConfig::new("obs", "https://mcp.example.com/sse") + .with_header("bad header", "value"); + let error = config.validate().expect_err("header names are validated"); + assert!(matches!(error, McpSseConfigError::InvalidHeaderName { .. })); +} + +#[test] +fn rejects_a_header_value_that_is_not_valid_for_http() { + let config = McpSseServerConfig::new("obs", "https://mcp.example.com/sse") + .with_header("authorization", "Bearer \nInjected: header"); + let error = config.validate().expect_err("header values are validated"); + assert!(matches!( + error, + McpSseConfigError::InvalidHeaderValue { .. } + )); +} + +#[test] +fn the_invalid_header_value_error_does_not_echo_the_value() { + let config = McpSseServerConfig::new("obs", "https://mcp.example.com/sse") + .with_header("authorization", "Bearer \nsuper-secret"); + let error = config.validate().expect_err("header values are validated"); + let rendered = error.to_string(); + assert!( + !rendered.contains("super-secret"), + "the error must not echo a secret: {rendered}" + ); + assert!(rendered.contains("authorization"), "got {rendered}"); +} + +// --------------------------------------------------------------------------- +// Defaults +// --------------------------------------------------------------------------- + +#[test] +fn default_limits_match_the_stdio_client_where_the_concept_is_shared() { + let limits = McpSseServerConfig::new("obs", "https://mcp.example.com/sse").limits; + assert_eq!( + limits.initialize_timeout, + std::time::Duration::from_secs(10) + ); + assert_eq!( + limits.list_tools_timeout, + std::time::Duration::from_secs(30) + ); + assert_eq!( + limits.call_tool_timeout, + std::time::Duration::from_secs(120) + ); +} + +#[test] +fn the_endpoint_limit_is_tighter_than_the_general_event_limit() { + let limits = McpSseServerConfig::new("obs", "https://mcp.example.com/sse").limits; + assert!( + limits.max_endpoint_bytes < limits.max_event_bytes, + "the pre-correlation surface must be smaller" + ); +} + +#[test] +fn a_config_deserializes_from_json_without_limits() { + let config: McpSseServerConfig = serde_json::from_value(serde_json::json!({ + "name": "obs", + "url": "https://mcp.example.com/sse", + "headers": {"authorization": "Bearer abc123"} + })) + .expect("deserialize"); + + assert_eq!(config.name, "obs"); + assert_eq!( + config + .headers + .get("authorization") + .expect("header") + .expose_secret(), + "Bearer abc123" + ); + assert_eq!(config.limits, super::McpSseLimits::default()); +} diff --git a/vendor/mentra/src/mcp/sse/endpoint.rs b/vendor/mentra/src/mcp/sse/endpoint.rs new file mode 100644 index 0000000..ae93870 --- /dev/null +++ b/vendor/mentra/src/mcp/sse/endpoint.rs @@ -0,0 +1,134 @@ +//! Endpoint URL resolution and same-origin enforcement. +//! +//! The legacy transport lets the *server* name the URL that the client will +//! POST JSON-RPC requests to, by sending it in an `endpoint` event. That makes +//! the endpoint value attacker-controlled whenever the server is compromised, +//! so it is validated before any request — and therefore any configured +//! `Authorization` header — is sent to it. +//! +//! The rule is deliberately strict: the resolved endpoint must share the +//! configured stream URL's scheme, host, and effective port. Anything else is +//! refused rather than normalized, because every relaxation here is a way to +//! redirect credentials to a host the operator never configured. + +#[cfg(test)] +mod tests; + +use url::Url; + +/// Schemes this transport is willing to speak. +const ALLOWED_SCHEMES: [&str; 2] = ["http", "https"]; + +/// Errors from validating a stream URL or a server-supplied endpoint. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum EndpointError { + #[error("the MCP server sent an empty endpoint event")] + Empty, + + #[error("could not parse the MCP endpoint URL: {0}")] + Malformed(String), + + #[error("unsupported MCP endpoint scheme '{scheme}': only http and https are allowed")] + UnsupportedScheme { scheme: String }, + + #[error("the MCP endpoint URL has no host")] + MissingHost, + + /// Rendered without the credentials themselves so a password in a + /// misconfigured URL never reaches a log. + #[error("the MCP endpoint URL must not embed credentials")] + CredentialsInUrl, + + #[error( + "the MCP server directed requests to '{endpoint}', which is not the configured origin '{configured}'" + )] + CrossOrigin { + endpoint: String, + configured: String, + }, +} + +/// Validates an operator-configured SSE stream URL. +/// +/// This runs before any connection is opened so that a bad configuration fails +/// at the boundary rather than mid-handshake. +pub(crate) fn validate_stream_url(raw: &str) -> Result { + let url = Url::parse(raw.trim()) + .map_err(|_| EndpointError::Malformed("invalid URL syntax".to_string()))?; + check_scheme(&url)?; + check_no_credentials(&url)?; + if url.host_str().is_none() { + return Err(EndpointError::MissingHost); + } + Ok(url) +} + +/// Resolves a server-supplied endpoint against the stream URL and enforces that +/// it stays on the same origin. +/// +/// `raw` is the `data` payload of the `endpoint` event. It is commonly a +/// relative path such as `/messages/?session_id=abc`, but the specification +/// also permits an absolute URL, so both are resolved through [`Url::join`]. +pub(crate) fn resolve_endpoint(stream_url: &Url, raw: &str) -> Result { + let trimmed = raw.trim(); + if trimmed.is_empty() { + return Err(EndpointError::Empty); + } + + let endpoint = stream_url + .join(trimmed) + .map_err(|_| EndpointError::Malformed("invalid URL syntax".to_string()))?; + + check_scheme(&endpoint)?; + check_no_credentials(&endpoint)?; + check_same_origin(stream_url, &endpoint)?; + + Ok(endpoint) +} + +/// Rejects any scheme outside the allowlist. +/// +/// This also covers `javascript:`, `data:`, and `file:`, which [`Url::join`] +/// happily produces from an absolute URL in the event payload. +fn check_scheme(url: &Url) -> Result<(), EndpointError> { + if ALLOWED_SCHEMES.contains(&url.scheme()) { + return Ok(()); + } + Err(EndpointError::UnsupportedScheme { + scheme: "[value omitted]".to_string(), + }) +} + +/// Rejects a URL carrying userinfo. +/// +/// [`Url::origin`] ignores userinfo, so without this check a server could send +/// `https://attacker@configured-host/` and pass the origin comparison while +/// changing what the client transmits. +fn check_no_credentials(url: &Url) -> Result<(), EndpointError> { + if url.username().is_empty() && url.password().is_none() { + return Ok(()); + } + Err(EndpointError::CredentialsInUrl) +} + +/// Requires an exact match on scheme, host, and effective port. +/// +/// The comparison uses [`Url::port_or_known_default`] so that an explicit +/// default port (`https://host:443`) and an implicit one (`https://host`) are +/// treated as the same origin, and [`Url::host`] rather than the raw string so +/// that equivalent IP literal spellings compare equal. Host names are compared +/// exactly: a trailing dot or a punycode homograph is a different origin. +fn check_same_origin(stream_url: &Url, endpoint: &Url) -> Result<(), EndpointError> { + let same = stream_url.scheme() == endpoint.scheme() + && stream_url.host() == endpoint.host() + && stream_url.port_or_known_default() == endpoint.port_or_known_default(); + + if same { + return Ok(()); + } + + Err(EndpointError::CrossOrigin { + endpoint: "[server-supplied origin omitted]".to_string(), + configured: "[configured origin]".to_string(), + }) +} diff --git a/vendor/mentra/src/mcp/sse/endpoint/tests.rs b/vendor/mentra/src/mcp/sse/endpoint/tests.rs new file mode 100644 index 0000000..5b8687a --- /dev/null +++ b/vendor/mentra/src/mcp/sse/endpoint/tests.rs @@ -0,0 +1,337 @@ +//! Tests for endpoint URL resolution and same-origin enforcement. + +use url::Url; + +use super::{EndpointError, resolve_endpoint, validate_stream_url}; + +/// The configured SSE URL every test resolves against. +const BASE: &str = "https://good-host.example/sse"; + +fn base() -> Url { + Url::parse(BASE).expect("base URL should parse") +} + +fn resolve(raw: &str) -> Result { + resolve_endpoint(&base(), raw) +} + +// --------------------------------------------------------------------------- +// Accepted endpoints +// --------------------------------------------------------------------------- + +#[test] +fn resolves_an_absolute_path_against_the_stream_url() { + let endpoint = resolve("/messages/?session_id=abc").expect("same-origin path is allowed"); + assert_eq!( + endpoint.as_str(), + "https://good-host.example/messages/?session_id=abc" + ); +} + +#[test] +fn resolves_a_relative_path_against_the_stream_url() { + let base = Url::parse("https://good-host.example/mcp/sse").expect("base should parse"); + let endpoint = resolve_endpoint(&base, "messages?session_id=abc").expect("relative is allowed"); + assert_eq!( + endpoint.as_str(), + "https://good-host.example/mcp/messages?session_id=abc" + ); +} + +#[test] +fn accepts_an_absolute_url_on_the_same_origin() { + let endpoint = + resolve("https://good-host.example/messages/").expect("same-origin absolute is allowed"); + assert_eq!(endpoint.as_str(), "https://good-host.example/messages/"); +} + +#[test] +fn accepts_an_explicit_default_port_matching_the_implicit_one() { + // https://host and https://host:443 are the same origin. + let endpoint = resolve("https://good-host.example:443/messages/") + .expect("default port is the same origin"); + assert_eq!(endpoint.as_str(), "https://good-host.example/messages/"); +} + +#[test] +fn accepts_an_implicit_default_port_matching_an_explicit_one() { + let base = Url::parse("https://good-host.example:443/sse").expect("base should parse"); + let endpoint = resolve_endpoint(&base, "https://good-host.example/messages/") + .expect("implicit port is the same origin"); + assert_eq!(endpoint.as_str(), "https://good-host.example/messages/"); +} + +#[test] +fn accepts_a_matching_non_default_port() { + let base = Url::parse("http://127.0.0.1:8080/sse").expect("base should parse"); + let endpoint = + resolve_endpoint(&base, "/messages/?session_id=abc").expect("same port is allowed"); + assert_eq!( + endpoint.as_str(), + "http://127.0.0.1:8080/messages/?session_id=abc" + ); +} + +#[test] +fn accepts_a_host_differing_only_by_case() { + // Host comparison is case-insensitive because the parser normalizes it. + let endpoint = + resolve("https://GOOD-HOST.EXAMPLE/messages/").expect("host case is not significant"); + assert_eq!(endpoint.as_str(), "https://good-host.example/messages/"); +} + +#[test] +fn accepts_a_plain_http_origin_when_the_stream_is_plain_http() { + let base = Url::parse("http://localhost:3000/sse").expect("base should parse"); + let endpoint = resolve_endpoint(&base, "/messages/").expect("http is an allowed scheme"); + assert_eq!(endpoint.as_str(), "http://localhost:3000/messages/"); +} + +#[test] +fn preserves_the_query_string_carrying_the_session_id() { + let endpoint = resolve("/messages/?session_id=6c8f2a&foo=bar").expect("query is preserved"); + assert_eq!(endpoint.query(), Some("session_id=6c8f2a&foo=bar")); +} + +// --------------------------------------------------------------------------- +// Rejected endpoints — cross-origin +// --------------------------------------------------------------------------- + +#[test] +fn rejects_an_absolute_url_on_a_different_host() { + let error = resolve("https://evil.example/steal").expect_err("cross-host must be rejected"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_protocol_relative_url_that_replaces_the_authority() { + // `//evil.example/x` inherits only the scheme; url::join gives it a NEW host. + let error = resolve("//evil.example/steal").expect_err("protocol-relative must be rejected"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_backslash_authority_that_url_normalizes_to_a_new_host() { + // url normalizes leading backslashes the way browsers do, yielding a new host. + let error = resolve("/\\evil.example/steal").expect_err("backslash authority must be rejected"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_scheme_downgrade_to_plain_http() { + let error = + resolve("http://good-host.example/messages/").expect_err("downgrade must be rejected"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_scheme_upgrade_to_https() { + let base = Url::parse("http://good-host.example/sse").expect("base should parse"); + let error = resolve_endpoint(&base, "https://good-host.example/messages/") + .expect_err("scheme change must be rejected"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_different_explicit_port() { + let error = + resolve("https://good-host.example:8443/messages/").expect_err("port change is rejected"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_trailing_dot_host_that_resolves_to_the_same_name() { + // `good-host.example.` is a distinct host string; treat it as cross-origin + // rather than guessing at DNS equivalence. + let error = + resolve("https://good-host.example./messages/").expect_err("trailing dot is rejected"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_punycode_homograph_host() { + let error = resolve("https://g\u{f6}\u{f6}d-host.example/messages/") + .expect_err("homograph host is rejected"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_subdomain_of_the_configured_host() { + let error = resolve("https://evil.good-host.example/messages/") + .expect_err("subdomains are a different origin"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +#[test] +fn rejects_a_suffix_extension_of_the_configured_host() { + let error = resolve("https://good-host.example.evil.test/messages/") + .expect_err("suffix extension is a different origin"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +// --------------------------------------------------------------------------- +// Rejected endpoints — credentials and schemes +// --------------------------------------------------------------------------- + +#[test] +fn rejects_userinfo_even_on_the_matching_origin() { + // url::Origin ignores userinfo, so an explicit check is required: credentials + // in the URL would be sent to the server alongside the configured headers. + let error = resolve("https://attacker@good-host.example/messages/") + .expect_err("userinfo must be rejected"); + assert!(matches!(error, EndpointError::CredentialsInUrl)); +} + +#[test] +fn rejects_a_password_in_the_endpoint_url() { + let error = resolve("https://user:secret@good-host.example/messages/") + .expect_err("password must be rejected"); + assert!(matches!(error, EndpointError::CredentialsInUrl)); +} + +#[test] +fn rejects_a_javascript_scheme() { + let error = resolve("javascript:alert(1)").expect_err("javascript must be rejected"); + assert!(matches!(error, EndpointError::UnsupportedScheme { .. })); +} + +#[test] +fn rejects_a_data_scheme() { + let error = resolve("data:text/plain,hi").expect_err("data must be rejected"); + assert!(matches!(error, EndpointError::UnsupportedScheme { .. })); +} + +#[test] +fn rejects_a_file_scheme() { + let error = resolve("file:///etc/passwd").expect_err("file must be rejected"); + assert!(matches!(error, EndpointError::UnsupportedScheme { .. })); +} + +// --------------------------------------------------------------------------- +// Rejected endpoints — malformed +// --------------------------------------------------------------------------- + +#[test] +fn rejects_an_empty_endpoint_payload() { + let error = resolve("").expect_err("an empty endpoint must be rejected"); + assert!(matches!(error, EndpointError::Empty)); +} + +#[test] +fn rejects_a_whitespace_only_endpoint_payload() { + let error = resolve(" ").expect_err("a blank endpoint must be rejected"); + assert!(matches!(error, EndpointError::Empty)); +} + +#[test] +fn rejects_an_unparseable_endpoint() { + let error = resolve("http://[not-an-address/x").expect_err("garbage must be rejected"); + assert!(matches!(error, EndpointError::Malformed(_))); +} + +#[test] +fn trims_surrounding_whitespace_before_resolving() { + // Servers occasionally pad the data field; trimming must happen before the + // origin check so it cannot be used to smuggle a different authority. + let endpoint = resolve(" /messages/?session_id=abc ").expect("padding is trimmed"); + assert_eq!( + endpoint.as_str(), + "https://good-host.example/messages/?session_id=abc" + ); +} + +#[test] +fn rejects_a_padded_cross_origin_endpoint() { + let error = + resolve(" https://evil.example/steal ").expect_err("padding does not bypass the check"); + assert!(matches!(error, EndpointError::CrossOrigin { .. })); +} + +// --------------------------------------------------------------------------- +// Stream URL validation +// --------------------------------------------------------------------------- + +#[test] +fn accepts_an_https_stream_url() { + let url = validate_stream_url("https://good-host.example/sse").expect("https is allowed"); + assert_eq!(url.scheme(), "https"); +} + +#[test] +fn accepts_an_http_stream_url() { + let url = validate_stream_url("http://127.0.0.1:9000/sse").expect("http is allowed"); + assert_eq!(url.scheme(), "http"); +} + +#[test] +fn rejects_a_stream_url_with_an_unsupported_scheme() { + let error = validate_stream_url("ws://good-host.example/sse").expect_err("ws is rejected"); + assert!(matches!(error, EndpointError::UnsupportedScheme { .. })); +} + +#[test] +fn rejects_a_stream_url_with_embedded_credentials() { + let error = validate_stream_url("https://user:pass@good-host.example/sse") + .expect_err("credentials are rejected"); + assert!(matches!(error, EndpointError::CredentialsInUrl)); +} + +#[test] +fn rejects_a_stream_url_without_a_host() { + let error = validate_stream_url("file:///tmp/sse").expect_err("a hostless URL is rejected"); + assert!(matches!( + error, + EndpointError::UnsupportedScheme { .. } | EndpointError::MissingHost + )); +} + +#[test] +fn rejects_an_unparseable_stream_url() { + let error = validate_stream_url("not a url").expect_err("garbage is rejected"); + assert!(matches!(error, EndpointError::Malformed(_))); +} + +// --------------------------------------------------------------------------- +// Error reporting +// --------------------------------------------------------------------------- + +#[test] +fn the_cross_origin_error_does_not_retain_either_origin() { + let error = + resolve("https://remote-canary.invalid/steal").expect_err("cross-origin is rejected"); + let rendered = error.to_string(); + let debug = format!("{error:?}"); + for origin in ["remote-canary.invalid", "good-host.example"] { + assert!(!rendered.contains(origin), "got {rendered}"); + assert!(!debug.contains(origin), "got {debug}"); + } +} + +#[test] +fn unsupported_scheme_errors_do_not_retain_the_scheme() { + let error = resolve("remote-canary:payload").expect_err("the scheme is unsupported"); + let rendered = error.to_string(); + let debug = format!("{error:?}"); + assert!(!rendered.contains("remote-canary"), "got {rendered}"); + assert!(!debug.contains("remote-canary"), "got {debug}"); +} + +#[test] +fn malformed_endpoint_errors_do_not_retain_the_payload() { + let error = resolve("http://[remote-canary.invalid").expect_err("the endpoint is malformed"); + let rendered = error.to_string(); + let debug = format!("{error:?}"); + assert!(!rendered.contains("remote-canary"), "got {rendered}"); + assert!(!debug.contains("remote-canary"), "got {debug}"); +} + +#[test] +fn the_credentials_error_does_not_echo_the_credentials() { + let error = resolve("https://user:hunter2@good-host.example/messages/") + .expect_err("credentials are rejected"); + let rendered = error.to_string(); + assert!( + !rendered.contains("hunter2"), + "the error must not echo a secret: {rendered}" + ); +} diff --git a/vendor/mentra/src/mcp/sse/testing.rs b/vendor/mentra/src/mcp/sse/testing.rs new file mode 100644 index 0000000..513d770 --- /dev/null +++ b/vendor/mentra/src/mcp/sse/testing.rs @@ -0,0 +1,497 @@ +//! A deterministic local HTTP+SSE server for transport tests. +//! +//! The transport needs a fixture that keeps one connection parked on a +//! long-lived `GET` while serving `POST` requests on other connections, and +//! that lets a test decide exactly when each SSE event reaches the client. A +//! raw [`TcpListener`] driven from [`std::thread`] gives that control with no +//! new dependencies, matching the fixtures already used in `mentra-provider`. +//! +//! Two properties keep the resulting tests deterministic rather than +//! timing-dependent: +//! +//! - **A thread per connection.** One accept loop hands each connection to its +//! own thread, so a blocking read on the parked `GET` cannot stop a `POST` +//! from being answered. +//! - **Chunked framing with a flush per event.** Chunked encoding is used +//! rather than read-until-close so that a clean end (`0\r\n\r\n`) and an +//! abrupt truncation are distinguishable by the client. Without the explicit +//! flush the operating system coalesces writes and the test would pass even +//! for a client that buffered the whole body. + +use std::collections::HashMap; +use std::io::{Read, Write}; +use std::net::{TcpListener, TcpStream}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Condvar, Mutex, mpsc}; +use std::thread; +use std::time::Duration; + +/// How a fixture answers one JSON-RPC `POST`. +#[derive(Debug, Clone)] +pub(crate) enum PostReply { + /// Answer `202 Accepted`, the reference servers' behavior. + Accepted, + /// Answer `200 OK` with an empty body. + Ok, + /// Answer with the given status and body. + Status { code: u16, body: String }, + /// Answer `307` pointing at another origin, which must not be followed. + Redirect { location: String }, + /// Close the connection without answering. + Drop, + /// Read the complete request, then withhold the response headers. + StallBeforeHeaders, + /// Send a successful response head, then withhold its declared body. + StallAfterHeaders, +} + +/// A request captured by the fixture. +#[derive(Debug, Clone)] +pub(crate) struct CapturedRequest { + pub(crate) method: String, + pub(crate) target: String, + pub(crate) headers: HashMap, + pub(crate) body: String, +} + +impl CapturedRequest { + /// Returns the JSON-RPC method of the captured body, if it has one. + pub(crate) fn rpc_method(&self) -> Option { + serde_json::from_str::(&self.body) + .ok()? + .get("method")? + .as_str() + .map(str::to_string) + } + + /// Returns the JSON-RPC id of the captured body, if it has one. + pub(crate) fn rpc_id(&self) -> Option { + serde_json::from_str::(&self.body) + .ok()? + .get("id")? + .as_u64() + } + + /// Returns a header value by its lowercase name. + pub(crate) fn header(&self, name: &str) -> Option<&str> { + self.headers.get(name).map(String::as_str) + } +} + +/// What the fixture should do when the SSE `GET` arrives. +#[derive(Debug, Clone)] +pub(crate) enum StreamOpening { + /// Accept the stream and serve events on demand. + Accept, + /// Answer `200` with a content type that is not `text/event-stream`. + WrongContentType, + /// Answer with the given status and body. + Status { code: u16, body: String }, + /// Answer `307` pointing elsewhere, which must not be followed. + Redirect { location: String }, +} + +/// Shared state between the test and the fixture's threads. +struct Shared { + requests: Mutex>, + replies: Mutex>, + stream_opened: (Mutex, Condvar), + posts_seen: (Mutex, Condvar), + post_headers_sent: (Mutex, Condvar), + stalled_posts_released: (Mutex, Condvar), + connections: AtomicUsize, +} + +/// Instruction sent to the thread holding the SSE connection open. +enum StreamCommand { + /// Write raw bytes as one chunk and flush. + Write(String), + /// End the stream cleanly with a terminal chunk. + Close, + /// Drop the connection without a terminal chunk, simulating a truncation. + Abort, +} + +/// A running local MCP HTTP+SSE server. +pub(crate) struct SseTestServer { + base_url: String, + shared: Arc, + commands: mpsc::Sender, +} + +impl SseTestServer { + /// Starts a fixture that accepts the SSE stream and answers every `POST` + /// with `202 Accepted`. + pub(crate) fn start() -> Self { + Self::with_opening(StreamOpening::Accept) + } + + /// Starts a fixture whose SSE `GET` is answered as described. + pub(crate) fn with_opening(opening: StreamOpening) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind the fixture listener"); + let addr = listener.local_addr().expect("read the fixture address"); + + let shared = Arc::new(Shared { + requests: Mutex::new(Vec::new()), + replies: Mutex::new(Vec::new()), + stream_opened: (Mutex::new(false), Condvar::new()), + posts_seen: (Mutex::new(0), Condvar::new()), + post_headers_sent: (Mutex::new(0), Condvar::new()), + stalled_posts_released: (Mutex::new(false), Condvar::new()), + connections: AtomicUsize::new(0), + }); + let (commands, command_rx) = mpsc::channel(); + + let accept_shared = Arc::clone(&shared); + let command_rx = Arc::new(Mutex::new(command_rx)); + thread::spawn(move || { + for incoming in listener.incoming() { + let Ok(stream) = incoming else { break }; + accept_shared.connections.fetch_add(1, Ordering::SeqCst); + let shared = Arc::clone(&accept_shared); + let command_rx = Arc::clone(&command_rx); + let opening = opening.clone(); + thread::spawn(move || serve_connection(stream, shared, command_rx, opening)); + } + }); + + Self { + base_url: format!("http://{addr}"), + shared, + commands, + } + } + + /// The URL of the SSE stream endpoint. + pub(crate) fn sse_url(&self) -> String { + format!("{}/sse", self.base_url) + } + + /// The fixture's origin, for building cross-origin cases. + pub(crate) fn base_url(&self) -> &str { + &self.base_url + } + + /// Queues the reply used for the next `POST`. + /// + /// Replies are consumed in order; once the queue is empty every `POST` is + /// answered `202 Accepted`. + pub(crate) fn queue_post_reply(&self, reply: PostReply) { + self.shared + .replies + .lock() + .expect("lock the reply queue") + .push(reply); + } + + /// Blocks until the client has opened the SSE stream. + pub(crate) fn wait_for_stream(&self) { + let (lock, condvar) = &self.shared.stream_opened; + let mut opened = lock.lock().expect("lock the stream flag"); + while !*opened { + let (guard, timeout) = condvar + .wait_timeout(opened, Duration::from_secs(10)) + .expect("wait for the stream"); + opened = guard; + assert!(!timeout.timed_out(), "the client never opened the stream"); + } + } + + /// Blocks until at least `count` `POST` requests have been captured. + /// + /// Tests synchronize on this rather than on a sleep so they stay + /// deterministic under load. + pub(crate) fn wait_for_posts(&self, count: usize) { + let (lock, condvar) = &self.shared.posts_seen; + let mut seen = lock.lock().expect("lock the post counter"); + while *seen < count { + let (guard, timeout) = condvar + .wait_timeout(seen, Duration::from_secs(10)) + .expect("wait for posts"); + seen = guard; + assert!( + !timeout.timed_out(), + "expected {count} POSTs, saw {}", + *lock.lock().expect("lock the post counter") + ); + } + } + + /// Blocks until at least `count` `POST` response heads have been flushed. + pub(crate) fn wait_for_post_response_headers(&self, count: usize) { + let (lock, condvar) = &self.shared.post_headers_sent; + let mut seen = lock.lock().expect("lock the POST response-head counter"); + while *seen < count { + let (guard, timeout) = condvar + .wait_timeout(seen, Duration::from_secs(10)) + .expect("wait for POST response headers"); + seen = guard; + assert!( + !timeout.timed_out(), + "expected {count} POST response heads, saw {}", + *lock.lock().expect("lock the POST response-head counter") + ); + } + } + + /// Releases every fixture connection deliberately stalled while replying. + pub(crate) fn release_stalled_posts(&self) { + let (lock, condvar) = &self.shared.stalled_posts_released; + *lock.lock().expect("lock the stalled-POST gate") = true; + condvar.notify_all(); + } + + /// Writes raw bytes to the SSE stream as one chunk. + pub(crate) fn send_raw(&self, payload: impl Into) { + let _ = self.commands.send(StreamCommand::Write(payload.into())); + } + + /// Writes an `endpoint` event naming the given POST target. + pub(crate) fn send_endpoint(&self, target: &str) { + self.send_raw(format!("event: endpoint\ndata: {target}\n\n")); + } + + /// Writes a `message` event carrying the given JSON-RPC payload. + pub(crate) fn send_message(&self, payload: &serde_json::Value) { + self.send_raw(format!("event: message\ndata: {payload}\n\n")); + } + + /// Ends the stream cleanly. + pub(crate) fn close_stream(&self) { + let _ = self.commands.send(StreamCommand::Close); + } + + /// Drops the stream connection without a terminal chunk. + pub(crate) fn abort_stream(&self) { + let _ = self.commands.send(StreamCommand::Abort); + } + + /// Every request the fixture has captured, in arrival order. + pub(crate) fn requests(&self) -> Vec { + self.shared + .requests + .lock() + .expect("lock the request log") + .clone() + } + + /// Only the `POST` requests captured so far. + pub(crate) fn posts(&self) -> Vec { + self.requests() + .into_iter() + .filter(|request| request.method == "POST") + .collect() + } +} + +/// Serves one accepted connection until the peer goes away. +fn serve_connection( + mut stream: TcpStream, + shared: Arc, + commands: Arc>>, + opening: StreamOpening, +) { + while let Some(request) = read_request(&mut stream) { + let is_stream_request = request.method == "GET"; + if !is_stream_request { + let (lock, condvar) = &shared.posts_seen; + let mut seen = lock.lock().expect("lock the post counter"); + *seen += 1; + condvar.notify_all(); + } + shared + .requests + .lock() + .expect("lock the request log") + .push(request); + + if is_stream_request { + serve_stream(&mut stream, &shared, &commands, &opening); + return; + } + + let reply = { + let mut replies = shared.replies.lock().expect("lock the reply queue"); + if replies.is_empty() { + PostReply::Accepted + } else { + replies.remove(0) + } + }; + + if !write_post_reply(&mut stream, reply, &shared) { + return; + } + } +} + +/// Answers the SSE `GET` and then streams events on command. +fn serve_stream( + stream: &mut TcpStream, + shared: &Arc, + commands: &Arc>>, + opening: &StreamOpening, +) { + let head = match opening { + StreamOpening::Accept => "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream; charset=utf-8\r\ntransfer-encoding: chunked\r\n\r\n".to_string(), + StreamOpening::WrongContentType => { + "HTTP/1.1 200 OK\r\ncontent-type: application/x-remote-canary\r\ncontent-length: 2\r\n\r\n{}".to_string() + } + StreamOpening::Status { code, body } => format!( + "HTTP/1.1 {code} Status\r\ncontent-type: text/plain\r\ncontent-length: {}\r\n\r\n{body}", + body.len() + ), + StreamOpening::Redirect { location } => format!( + "HTTP/1.1 307 Temporary Redirect\r\nlocation: {location}\r\ncontent-length: 0\r\n\r\n" + ), + }; + + if stream.write_all(head.as_bytes()).is_err() { + return; + } + let _ = stream.flush(); + + if !matches!(opening, StreamOpening::Accept) { + return; + } + + // Only signal readiness once the client can actually receive events. + let (lock, condvar) = &shared.stream_opened; + *lock.lock().expect("lock the stream flag") = true; + condvar.notify_all(); + + let commands = commands.lock().expect("lock the command channel"); + while let Ok(command) = commands.recv() { + match command { + StreamCommand::Write(payload) => { + let framed = format!("{:X}\r\n{payload}\r\n", payload.len()); + if stream.write_all(framed.as_bytes()).is_err() { + return; + } + // Flush per event, or the OS coalesces writes and the test + // would pass for a client that buffered the whole body. + let _ = stream.flush(); + } + StreamCommand::Close => { + let _ = stream.write_all(b"0\r\n\r\n"); + let _ = stream.flush(); + return; + } + StreamCommand::Abort => return, + } + } +} + +/// Writes one `POST` reply, reporting whether the connection may be reused. +fn write_post_reply(stream: &mut TcpStream, reply: PostReply, shared: &Shared) -> bool { + let stall_before_headers = matches!(&reply, PostReply::StallBeforeHeaders); + let stall_after_headers = matches!(&reply, PostReply::StallAfterHeaders); + + if stall_before_headers { + wait_for_stalled_post_release(shared); + return false; + } + + let response = match reply { + PostReply::Accepted => "HTTP/1.1 202 Accepted\r\ncontent-length: 0\r\n\r\n".to_string(), + PostReply::Ok => "HTTP/1.1 200 OK\r\ncontent-length: 0\r\n\r\n".to_string(), + PostReply::Status { code, body } => format!( + "HTTP/1.1 {code} Status\r\ncontent-type: text/plain\r\ncontent-length: {}\r\n\r\n{body}", + body.len() + ), + PostReply::Redirect { location } => format!( + "HTTP/1.1 307 Temporary Redirect\r\nlocation: {location}\r\ncontent-length: 0\r\n\r\n" + ), + PostReply::Drop => return false, + PostReply::StallAfterHeaders => "HTTP/1.1 200 OK\r\ncontent-length: 4\r\n\r\n".to_string(), + PostReply::StallBeforeHeaders => unreachable!("handled before building the response"), + }; + + if stream.write_all(response.as_bytes()).is_err() { + return false; + } + if stream.flush().is_err() { + return false; + } + + let (lock, condvar) = &shared.post_headers_sent; + *lock.lock().expect("lock the POST response-head counter") += 1; + condvar.notify_all(); + + if stall_after_headers { + wait_for_stalled_post_release(shared); + return false; + } + + true +} + +/// Waits until the test explicitly releases a deliberately stalled reply. +fn wait_for_stalled_post_release(shared: &Shared) { + let (lock, condvar) = &shared.stalled_posts_released; + let mut released = lock.lock().expect("lock the stalled-POST gate"); + while !*released { + released = condvar + .wait(released) + .expect("wait for the stalled POST to be released"); + } +} + +/// Reads one HTTP request, honoring keep-alive by returning `None` at EOF. +fn read_request(stream: &mut TcpStream) -> Option { + let mut buffer = Vec::new(); + let mut chunk = [0_u8; 1024]; + let mut header_end = None; + let mut content_length = 0_usize; + + loop { + let read = match stream.read(&mut chunk) { + Ok(0) | Err(_) => return None, + Ok(read) => read, + }; + buffer.extend_from_slice(&chunk[..read]); + + if header_end.is_none() { + let Some(index) = buffer.windows(4).position(|window| window == b"\r\n\r\n") else { + continue; + }; + let end = index + 4; + header_end = Some(end); + content_length = String::from_utf8_lossy(&buffer[..end]) + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().unwrap_or_default()) + }) + .unwrap_or_default(); + } + + if header_end.is_some_and(|end| buffer.len() >= end + content_length) { + break; + } + } + + let end = header_end?; + let head = String::from_utf8_lossy(&buffer[..end]).to_string(); + let body = String::from_utf8_lossy(&buffer[end..end + content_length]).to_string(); + + let mut lines = head.lines(); + let mut request_line = lines.next()?.split_whitespace(); + let method = request_line.next()?.to_string(); + let target = request_line.next()?.to_string(); + + let headers = lines + .filter_map(|line| { + let (name, value) = line.split_once(':')?; + Some((name.trim().to_ascii_lowercase(), value.trim().to_string())) + }) + .collect(); + + Some(CapturedRequest { + method, + target, + headers, + body, + }) +} diff --git a/vendor/mentra/src/mcp/sse/wire.rs b/vendor/mentra/src/mcp/sse/wire.rs new file mode 100644 index 0000000..8c96c31 --- /dev/null +++ b/vendor/mentra/src/mcp/sse/wire.rs @@ -0,0 +1,207 @@ +//! Incremental Server-Sent Events wire parser. +//! +//! This is a byte-oriented push parser: callers feed arbitrary byte chunks as +//! they arrive from the transport and receive whole dispatched events back. It +//! deliberately buffers bytes rather than strings so that a UTF-8 sequence or a +//! CRLF pair split across two network chunks is reassembled correctly. +//! +//! The framing rules follow the WHATWG `text/event-stream` interpretation: +//! +//! - lines end with `\n`, `\r\n`, or a lone `\r`; +//! - a line beginning with `:` is a comment (servers use these as heartbeats); +//! - `field: value` strips at most one space after the colon; +//! - a line with no colon is a field name with an empty value; +//! - repeated `data` fields are joined with `\n`; +//! - a blank line dispatches the buffered event, and an event with no `data` +//! field is discarded rather than dispatched; +//! - a leading UTF-8 byte order mark is ignored. +//! +//! Every buffered event is bounded by a caller-supplied limit so a hostile or +//! malfunctioning server cannot force unbounded memory growth. Mentra's tool +//! result limiter runs far too late to protect this parser. + +#[cfg(test)] +mod tests; + +/// Byte order mark that may prefix the very first line of a stream. +const UTF8_BOM: &str = "\u{feff}"; + +/// A dispatched Server-Sent Event. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct SseEvent { + /// The `event:` field, defaulting to `message` when the server omits it. + pub(crate) event: String, + /// The joined `data:` field values, without the trailing newline. + pub(crate) data: String, +} + +/// Errors produced while decoding the SSE byte stream. +/// +/// Public because it is reachable through +/// [`McpSseError::Wire`](crate::mcp::McpSseError::Wire); the parser itself +/// stays internal. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum SseWireError { + #[error("SSE event exceeded the {limit} byte limit (buffered at least {observed} bytes)")] + EventTooLarge { limit: usize, observed: usize }, + + #[error("SSE stream contained invalid UTF-8")] + InvalidUtf8, +} + +/// Incremental parser over the `text/event-stream` framing. +#[derive(Debug)] +pub(crate) struct SseParser { + /// Bytes of the line currently being accumulated. + line: Vec, + /// Joined `data:` values for the event currently being accumulated. + data: String, + /// The `event:` value for the event currently being accumulated. + event: Option, + /// Whether the previous byte was a carriage return whose companion line + /// feed may still arrive in a later chunk. + pending_cr: bool, + /// Whether the next completed line is the first of the stream and may + /// therefore carry a byte order mark. + at_stream_start: bool, + /// Maximum bytes buffered for a single event. + max_event_bytes: usize, +} + +impl SseParser { + /// Creates a parser that rejects any single event larger than + /// `max_event_bytes`. + pub(crate) fn new(max_event_bytes: usize) -> Self { + Self { + line: Vec::new(), + data: String::new(), + event: None, + pending_cr: false, + at_stream_start: true, + max_event_bytes, + } + } + + /// Feeds the next chunk of stream bytes and returns every event completed + /// by it. + /// + /// An error leaves the parser poisoned by contract: the caller must tear + /// the stream down rather than continue feeding it, because a size or + /// encoding violation means the framing can no longer be trusted. + pub(crate) fn feed(&mut self, bytes: &[u8]) -> Result, SseWireError> { + let mut events = Vec::new(); + + for byte in bytes { + let byte = *byte; + + // A carriage return already ended the previous line. If its + // companion line feed arrives now, it is part of that same + // terminator and must not end an additional (empty) line. + if self.pending_cr { + self.pending_cr = false; + if byte == b'\n' { + continue; + } + } + + match byte { + b'\n' => self.end_line(&mut events)?, + b'\r' => { + self.pending_cr = true; + self.end_line(&mut events)?; + } + _ => { + self.line.push(byte); + self.check_bounds()?; + } + } + } + + Ok(events) + } + + /// Rejects an event whose buffered bytes exceed the configured limit. + fn check_bounds(&self) -> Result<(), SseWireError> { + let observed = self + .event + .as_ref() + .map_or(0, String::len) + .saturating_add(self.data.len()) + .saturating_add(self.line.len()); + if observed > self.max_event_bytes { + return Err(SseWireError::EventTooLarge { + limit: self.max_event_bytes, + observed, + }); + } + Ok(()) + } + + /// Consumes the accumulated line, applying it to the pending event. + fn end_line(&mut self, events: &mut Vec) -> Result<(), SseWireError> { + let line = std::mem::take(&mut self.line); + let line = std::str::from_utf8(&line).map_err(|_| SseWireError::InvalidUtf8)?; + + // Only the very first line of the stream may carry a byte order mark. + let line = if std::mem::take(&mut self.at_stream_start) { + line.strip_prefix(UTF8_BOM).unwrap_or(line) + } else { + line + }; + + // A blank line dispatches whatever has been accumulated. + if line.is_empty() { + if let Some(event) = self.take_event() { + events.push(event); + } + return Ok(()); + } + + // A leading colon marks a comment, which servers use as a heartbeat. + if line.starts_with(':') { + return Ok(()); + } + + let (field, value) = match line.split_once(':') { + Some((field, value)) => (field, value.strip_prefix(' ').unwrap_or(value)), + // A line with no colon is a field name with an empty value. + None => (line, ""), + }; + + match field { + "event" => self.event = Some(value.to_string()), + "data" => { + self.data.push_str(value); + self.data.push('\n'); + self.check_bounds()?; + } + // `id` and `retry` belong to reconnection, which this transport + // does not implement; every other field is undefined and ignored. + _ => {} + } + + Ok(()) + } + + /// Takes the accumulated event, resetting the per-event state. + /// + /// Returns `None` when no `data` field was seen, which the specification + /// requires be discarded rather than dispatched. + fn take_event(&mut self) -> Option { + let event = self.event.take(); + let mut data = std::mem::take(&mut self.data); + + if data.is_empty() { + return None; + } + + // The dispatch step drops the single trailing newline added by the + // last `data` field. + data.pop(); + + Some(SseEvent { + event: event.unwrap_or_else(|| "message".to_string()), + data, + }) + } +} diff --git a/vendor/mentra/src/mcp/sse/wire/tests.rs b/vendor/mentra/src/mcp/sse/wire/tests.rs new file mode 100644 index 0000000..cb9e420 --- /dev/null +++ b/vendor/mentra/src/mcp/sse/wire/tests.rs @@ -0,0 +1,334 @@ +//! Tests for the incremental Server-Sent Events wire parser. + +use super::{SseEvent, SseParser, SseWireError}; + +/// Feed a whole payload as one chunk and collect every dispatched event. +fn parse_all(payload: &str) -> Vec { + let mut parser = SseParser::new(64 * 1024); + parser + .feed(payload.as_bytes()) + .expect("payload should parse") +} + +#[test] +fn dispatches_a_simple_message_event() { + let events = parse_all("event: message\ndata: hello\n\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].event, "message"); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn defaults_the_event_name_to_message_when_absent() { + let events = parse_all("data: hello\n\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].event, "message"); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn parses_the_endpoint_event_name() { + let events = parse_all("event: endpoint\ndata: /messages/?session_id=abc\n\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].event, "endpoint"); + assert_eq!(events[0].data, "/messages/?session_id=abc"); +} + +#[test] +fn accepts_crlf_line_terminators() { + let events = parse_all("event: message\r\ndata: hello\r\n\r\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].event, "message"); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn accepts_lone_cr_line_terminators() { + let events = parse_all("event: message\rdata: hello\r\r"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].event, "message"); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn joins_multiple_data_lines_with_newlines() { + let events = parse_all("event: message\ndata: first\ndata: second\ndata: third\n\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].data, "first\nsecond\nthird"); +} + +#[test] +fn preserves_embedded_json_across_multiple_data_lines() { + let events = parse_all("event: message\ndata: {\"jsonrpc\":\"2.0\",\ndata: \"id\":1}\n\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].data, "{\"jsonrpc\":\"2.0\",\n\"id\":1}"); +} + +#[test] +fn ignores_comment_and_heartbeat_lines() { + let events = parse_all(": ping\n: keep-alive\nevent: message\ndata: hello\n\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn ignores_a_standalone_heartbeat_without_dispatching() { + let events = parse_all(": heartbeat\n\n"); + assert!(events.is_empty()); +} + +#[test] +fn strips_only_one_leading_space_from_a_field_value() { + let events = parse_all("data: two-spaces\n\n"); + assert_eq!(events[0].data, " two-spaces"); +} + +#[test] +fn accepts_a_data_field_with_no_space_after_the_colon() { + let events = parse_all("data:hello\n\n"); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn treats_a_bare_field_name_as_an_empty_value() { + let events = parse_all("data\ndata: hello\n\n"); + assert_eq!(events[0].data, "\nhello"); +} + +#[test] +fn ignores_unknown_fields() { + let events = parse_all("id: 42\nretry: 3000\nfoo: bar\ndata: hello\n\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn does_not_dispatch_an_event_without_data() { + let events = parse_all("event: message\n\n"); + assert!(events.is_empty()); +} + +#[test] +fn resets_the_event_name_between_dispatches() { + let events = parse_all("event: endpoint\ndata: /messages\n\ndata: hello\n\n"); + assert_eq!(events.len(), 2); + assert_eq!(events[0].event, "endpoint"); + assert_eq!(events[1].event, "message"); +} + +#[test] +fn dispatches_several_events_from_one_chunk() { + let events = parse_all("data: one\n\ndata: two\n\ndata: three\n\n"); + assert_eq!(events.len(), 3); + assert_eq!(events[0].data, "one"); + assert_eq!(events[1].data, "two"); + assert_eq!(events[2].data, "three"); +} + +#[test] +fn strips_a_leading_utf8_byte_order_mark() { + let mut parser = SseParser::new(64 * 1024); + let mut payload = vec![0xEF, 0xBB, 0xBF]; + payload.extend_from_slice(b"data: hello\n\n"); + let events = parser.feed(&payload).expect("payload should parse"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn reassembles_an_event_split_across_arbitrary_chunks() { + let payload = "event: message\ndata: {\"jsonrpc\":\"2.0\",\"id\":1}\n\n"; + let bytes = payload.as_bytes(); + // Split at every possible byte boundary; the parse must be identical. + for split in 1..bytes.len() { + let mut parser = SseParser::new(64 * 1024); + let mut events = parser + .feed(&bytes[..split]) + .expect("first half should parse"); + events.extend( + parser + .feed(&bytes[split..]) + .expect("second half should parse"), + ); + assert_eq!( + events.len(), + 1, + "split at {split} should dispatch one event" + ); + assert_eq!(events[0].event, "message"); + assert_eq!(events[0].data, "{\"jsonrpc\":\"2.0\",\"id\":1}"); + } +} + +#[test] +fn reassembles_a_crlf_event_split_between_the_cr_and_the_lf() { + let payload = "data: hello\r\n\r\n"; + let bytes = payload.as_bytes(); + let split = payload.find('\r').expect("payload has a CR") + 1; + let mut parser = SseParser::new(64 * 1024); + let mut events = parser + .feed(&bytes[..split]) + .expect("first half should parse"); + assert!(events.is_empty(), "a dangling CR must not dispatch yet"); + events.extend( + parser + .feed(&bytes[split..]) + .expect("second half should parse"), + ); + assert_eq!(events.len(), 1); + assert_eq!(events[0].data, "hello"); +} + +#[test] +fn reassembles_a_multibyte_character_split_across_chunks() { + let payload = "data: caf\u{e9}\n\n"; + let bytes = payload.as_bytes(); + // The 'é' is two bytes; split between them. + let split = bytes + .iter() + .position(|byte| *byte == 0xC3) + .expect("payload has a two-byte character") + + 1; + let mut parser = SseParser::new(64 * 1024); + let mut events = parser + .feed(&bytes[..split]) + .expect("first half should parse"); + events.extend( + parser + .feed(&bytes[split..]) + .expect("second half should parse"), + ); + assert_eq!(events.len(), 1); + assert_eq!(events[0].data, "caf\u{e9}"); +} + +#[test] +fn feeding_one_byte_at_a_time_matches_a_single_chunk() { + let payload = "event: endpoint\r\ndata: /messages/?session_id=abc\r\n\r\ndata: tail\n\n"; + let mut parser = SseParser::new(64 * 1024); + let mut events = Vec::new(); + for byte in payload.as_bytes() { + events.extend( + parser + .feed(std::slice::from_ref(byte)) + .expect("byte should parse"), + ); + } + assert_eq!(events.len(), 2); + assert_eq!(events[0].event, "endpoint"); + assert_eq!(events[0].data, "/messages/?session_id=abc"); + assert_eq!(events[1].event, "message"); + assert_eq!(events[1].data, "tail"); +} + +#[test] +fn rejects_an_event_larger_than_the_configured_limit() { + let mut parser = SseParser::new(64); + let oversized = format!("data: {}\n\n", "x".repeat(512)); + let error = parser + .feed(oversized.as_bytes()) + .expect_err("oversized event must be rejected"); + assert!(matches!( + error, + SseWireError::EventTooLarge { limit: 64, .. } + )); +} + +#[test] +fn rejects_an_unterminated_line_larger_than_the_configured_limit() { + let mut parser = SseParser::new(64); + // No terminator at all: the parser must not buffer without bound. + let error = parser + .feed("x".repeat(512).as_bytes()) + .expect_err("oversized line must be rejected"); + assert!(matches!( + error, + SseWireError::EventTooLarge { limit: 64, .. } + )); +} + +#[test] +fn rejects_an_event_that_only_exceeds_the_limit_across_several_data_lines() { + let mut parser = SseParser::new(128); + let line = format!("data: {}\n", "x".repeat(60)); + let mut error = None; + for _ in 0..10 { + if let Err(e) = parser.feed(line.as_bytes()) { + error = Some(e); + break; + } + } + assert!( + matches!(error, Some(SseWireError::EventTooLarge { limit: 128, .. })), + "accumulated data across lines must be bounded, got {error:?}" + ); +} + +#[test] +fn counts_the_stored_event_name_toward_the_event_limit() { + let mut parser = SseParser::new(64); + let event_name = "x".repeat(57); + + parser + .feed(format!("event: {event_name}\n").as_bytes()) + .expect("the event line exactly fills the limit"); + let error = parser + .feed(b"data: xx") + .expect_err("stored event name and current data line must share the limit"); + + assert_eq!( + error, + SseWireError::EventTooLarge { + limit: 64, + observed: 65, + } + ); +} + +#[test] +fn accepts_an_event_whose_total_buffered_size_exactly_matches_the_limit() { + let mut parser = SseParser::new(64); + let event_name = "x".repeat(57); + let payload = format!("event: {event_name}\ndata: x\n\n"); + + let events = parser + .feed(payload.as_bytes()) + .expect("an event at the exact byte limit should parse"); + + assert_eq!( + events, + vec![SseEvent { + event: event_name, + data: "x".to_string(), + }] + ); +} + +#[test] +fn accounts_size_per_event_rather_than_per_stream() { + let mut parser = SseParser::new(64); + // Each event is small; many of them in sequence must not trip the limit. + for _ in 0..50 { + let events = parser + .feed(b"data: small\n\n") + .expect("each small event should parse"); + assert_eq!(events.len(), 1); + } +} + +#[test] +fn rejects_invalid_utf8_in_the_stream() { + let mut parser = SseParser::new(64 * 1024); + let error = parser + .feed(&[b'd', b'a', b't', b'a', b':', b' ', 0xFF, 0xFE, b'\n', b'\n']) + .expect_err("invalid UTF-8 must be rejected"); + assert!(matches!(error, SseWireError::InvalidUtf8)); +} + +#[test] +fn does_not_dispatch_a_trailing_event_without_a_blank_line() { + // A stream that ends mid-event must not yield a truncated event. + let events = parse_all("data: complete\n\ndata: incomplete\n"); + assert_eq!(events.len(), 1); + assert_eq!(events[0].data, "complete"); +} diff --git a/vendor/mentra/src/mcp/tests.rs b/vendor/mentra/src/mcp/tests.rs new file mode 100644 index 0000000..b7e0e5d --- /dev/null +++ b/vendor/mentra/src/mcp/tests.rs @@ -0,0 +1,298 @@ +use std::{collections::BTreeMap, sync::Arc}; + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::{ + ContentBlock, + mcp::{ + bridge::{McpBridgedTool, McpToolClient, mcp_tool_name, parse_mcp_tool_name}, + protocol::*, + }, + runtime::RuntimePolicy, + test::{MockRuntime, MockToolCall}, + tool::{ + ParallelToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSideEffectLevel, + ToolSpec, + }, +}; + +pub(crate) struct SuccessfulMcpClient { + pub(crate) output: String, +} + +#[async_trait] +impl McpToolClient for SuccessfulMcpClient { + async fn call_tool( + &self, + _tool_name: &str, + _arguments: Option, + ) -> Result { + Ok(McpToolCallResult { + content: vec![McpToolCallContent { + kind: "text".to_string(), + text: Some(self.output.clone()), + data: None, + mime_type: None, + }], + is_error: false, + }) + } +} + +struct MatchingCustomTool { + output: String, +} + +impl ToolDefinition for MatchingCustomTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("matching_custom_output") + .description("Return the same output as the MCP test tool") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::External) + .build() + } +} + +#[async_trait] +impl ToolExecutor for MatchingCustomTool { + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + Ok(self.output.clone()) + } +} + +#[test] +fn mcp_tool_name_namespacing() { + assert_eq!(mcp_tool_name("filesystem", "read"), "mcp__filesystem__read"); + assert_eq!( + mcp_tool_name("my-server", "do_thing"), + "mcp__my-server__do_thing" + ); +} + +/// The bridge accepts either transport's client without a signature change at +/// the call site, which is what keeps `McpBridgedTool::new` source compatible +/// for existing stdio callers. +#[test] +fn bridging_compiles_for_both_transport_clients() { + fn accepts_stdio(client: Arc) -> McpBridgedTool { + McpBridgedTool::new( + "stdio-server".to_string(), + McpToolDefinition { + name: "read".to_string(), + description: None, + input_schema: None, + }, + client, + ) + } + + fn accepts_sse(client: Arc) -> McpBridgedTool { + McpBridgedTool::new( + "sse-server".to_string(), + McpToolDefinition { + name: "search".to_string(), + description: None, + input_schema: None, + }, + client, + ) + } + + // Building a real client of either transport needs a live server, so this + // asserts the signatures rather than the behavior. + let _ = accepts_stdio; + let _ = accepts_sse; +} + +#[test] +fn parse_mcp_tool_name_roundtrip() { + let name = mcp_tool_name("filesystem", "read"); + let (server, tool) = parse_mcp_tool_name(&name).expect("should parse"); + assert_eq!(server, "filesystem"); + assert_eq!(tool, "read"); +} + +#[test] +fn parse_mcp_tool_name_rejects_non_mcp() { + assert!(parse_mcp_tool_name("regular_tool").is_none()); + assert!(parse_mcp_tool_name("mcp_no_double_underscore").is_none()); +} + +#[tokio::test] +async fn bridged_output_is_truncated_before_the_next_provider_request() { + let full_output = "one\ntwo\nthree"; + let bridged_name = mcp_tool_name("fake", "large_output"); + let mock = MockRuntime::builder() + .with_policy( + RuntimePolicy::permissive() + .with_max_tool_result_bytes(8) + .with_max_tool_result_lines(1) + .spill_full_tool_output(false), + ) + .tool_calls([ + MockToolCall::new(&bridged_name, json!({})).with_id("mcp-call"), + MockToolCall::new("matching_custom_output", json!({})).with_id("custom-call"), + ]) + .text("done") + .build() + .expect("build mock runtime"); + mock.runtime().register_tool(McpBridgedTool::new_for_test( + "fake".to_string(), + McpToolDefinition { + name: "large_output".to_string(), + description: Some("Return an oversized result".to_string()), + input_schema: Some(json!({ "type": "object", "properties": {} })), + }, + Arc::new(SuccessfulMcpClient { + output: full_output.to_string(), + }), + )); + mock.runtime().register_tool(MatchingCustomTool { + output: full_output.to_string(), + }); + let mut agent = mock + .runtime() + .spawn("mcp-truncation-test", mock.model()) + .expect("spawn agent"); + + let response = agent + .send(vec![ContentBlock::text("run both tools")]) + .await + .expect("run agent"); + assert_eq!(response.text(), "done"); + + let requests = mock.recorded_requests().await; + assert_eq!(requests.len(), 2); + let provider_results = requests[1] + .messages + .iter() + .flat_map(|message| &message.content) + .filter_map(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => Some((tool_use_id.as_str(), (content.as_str(), *is_error))), + _ => None, + }) + .collect::>(); + let mcp_result = provider_results + .get("mcp-call") + .expect("provider request should contain the MCP result"); + let custom_result = provider_results + .get("custom-call") + .expect("provider request should contain the custom-tool result"); + assert_eq!(mcp_result, custom_result); + assert_eq!( + *mcp_result, + ( + "one\n[truncated: showing 1 of 3 lines; full output was not saved because spill-to-file is disabled by runtime policy]", + false, + ) + ); +} + +#[test] +fn json_rpc_request_serialization() { + let req = JsonRpcRequest::new(1, "initialize", Some(json!({"key": "value"}))); + let serialized = serde_json::to_string(&req).expect("serialize"); + assert!(serialized.contains("\"jsonrpc\":\"2.0\"")); + assert!(serialized.contains("\"id\":1")); + assert!(serialized.contains("\"method\":\"initialize\"")); +} + +#[test] +fn json_rpc_response_deserialization() { + let json = r#"{"jsonrpc":"2.0","id":1,"result":{"tools":[]}}"#; + let resp: JsonRpcResponse = serde_json::from_str(json).expect("deserialize"); + assert_eq!(resp.id, JsonRpcId::Number(1)); + assert!(resp.result.is_some()); + assert!(resp.error.is_none()); +} + +#[test] +fn json_rpc_error_response_deserialization() { + let json = r#"{"jsonrpc":"2.0","id":2,"error":{"code":-32600,"message":"Invalid Request"}}"#; + let resp: JsonRpcResponse = serde_json::from_str(json).expect("deserialize"); + assert_eq!(resp.id, JsonRpcId::Number(2)); + let err = resp.error.expect("should have error"); + assert_eq!(err.code, -32600); + assert_eq!(err.message, "Invalid Request"); +} + +#[test] +fn mcp_tool_definition_deserialization() { + let json = json!({ + "name": "read_file", + "description": "Read a file from disk", + "inputSchema": { + "type": "object", + "properties": { + "path": {"type": "string"} + }, + "required": ["path"] + } + }); + let tool: McpToolDefinition = serde_json::from_value(json).expect("deserialize"); + assert_eq!(tool.name, "read_file"); + assert_eq!(tool.description.as_deref(), Some("Read a file from disk")); + assert!(tool.input_schema.is_some()); +} + +#[test] +fn mcp_tool_call_result_deserialization() { + let json = json!({ + "content": [ + {"type": "text", "text": "Hello, world!"}, + {"type": "text", "text": "Second block"} + ], + "isError": false + }); + let result: McpToolCallResult = serde_json::from_value(json).expect("deserialize"); + assert_eq!(result.content.len(), 2); + assert!(!result.is_error); + assert_eq!(result.content[0].text.as_deref(), Some("Hello, world!")); +} + +#[test] +fn mcp_tool_call_error_result() { + let json = json!({ + "content": [{"type": "text", "text": "Something went wrong"}], + "isError": true + }); + let result: McpToolCallResult = serde_json::from_value(json).expect("deserialize"); + assert!(result.is_error); +} + +#[test] +fn mcp_server_config_deserialization() { + let json = json!({ + "name": "filesystem", + "command": "npx", + "args": ["-y", "@modelcontextprotocol/server-filesystem", "/tmp"], + "env": {"DEBUG": "1"}, + "cwd": "/home/user" + }); + let config: McpServerConfig = serde_json::from_value(json).expect("deserialize"); + assert_eq!(config.name, "filesystem"); + assert_eq!(config.command, "npx"); + assert_eq!(config.args.len(), 3); + assert_eq!(config.env.get("DEBUG").map(String::as_str), Some("1")); + assert_eq!(config.cwd.as_deref(), Some("/home/user")); +} + +#[test] +fn mcp_initialize_params_serialization() { + let params = McpInitializeParams { + protocol_version: "2024-11-05".to_string(), + capabilities: json!({}), + client_info: McpClientInfo { + name: "mentra".to_string(), + version: "0.6.0".to_string(), + }, + }; + let json = serde_json::to_value(¶ms).expect("serialize"); + assert_eq!(json["protocolVersion"], "2024-11-05"); + assert_eq!(json["clientInfo"]["name"], "mentra"); +} diff --git a/vendor/mentra/src/memory.rs b/vendor/mentra/src/memory.rs new file mode 100644 index 0000000..6b109a7 --- /dev/null +++ b/vendor/mentra/src/memory.rs @@ -0,0 +1,16 @@ +mod compaction; +mod engine; +mod hybrid_store; +pub(crate) mod journal; + +pub(crate) use compaction::{ + estimated_request_tokens, micro_compact_history, required_tail_start_for_continuation, +}; +pub use engine::{ + IngestOutcome, IngestRequest, MAX_MEMORY_LIST_PAGE_SIZE, MemoryCursor, MemoryEngine, MemoryHit, + MemoryListCursor, MemoryListFilter, MemoryListPage, MemoryListRequest, MemoryListSort, + MemoryRecord, MemoryRecordKind, MemorySearchMode, MemorySearchRequest, MemoryStore, + SearchRequest, +}; +pub(crate) use engine::{build_search_query, recalled_memory_message}; +pub use hybrid_store::SqliteHybridMemoryStore; diff --git a/vendor/mentra/src/memory/compaction.rs b/vendor/mentra/src/memory/compaction.rs new file mode 100644 index 0000000..b434058 --- /dev/null +++ b/vendor/mentra/src/memory/compaction.rs @@ -0,0 +1,107 @@ +use std::collections::HashMap; + +use crate::{ContentBlock, Message, Role}; + +const MICRO_COMPACT_MIN_CONTENT_LEN: usize = 100; + +pub(crate) fn micro_compact_history(history: &[Message], keep_recent: usize) -> Vec { + if keep_recent == usize::MAX { + return history.to_vec(); + } + + let mut compacted = history.to_vec(); + let tool_names = tool_name_index(&compacted); + let mut tool_results = Vec::new(); + + for (message_index, message) in compacted.iter().enumerate() { + if message.role != Role::User { + continue; + } + + for (block_index, block) in message.content.iter().enumerate() { + if matches!(block, ContentBlock::ToolResult { .. }) { + tool_results.push((message_index, block_index)); + } + } + } + + if tool_results.len() <= keep_recent { + return compacted; + } + + let compact_count = tool_results.len() - keep_recent; + for (message_index, block_index) in tool_results.into_iter().take(compact_count) { + let Some(ContentBlock::ToolResult { + tool_use_id, + content, + .. + }) = compacted[message_index].content.get_mut(block_index) + else { + continue; + }; + + if content.len() <= MICRO_COMPACT_MIN_CONTENT_LEN { + continue; + } + + let tool_name = tool_names + .get(tool_use_id.as_str()) + .map(String::as_str) + .unwrap_or("tool"); + content.clear(); + content.push_str(&format!("[Previous: used {tool_name}]")); + } + + compacted +} + +pub(crate) fn estimated_request_tokens(messages: &[Message], system: Option<&str>) -> usize { + let mut estimated = + estimated_tokens_for_str(&serde_json::to_string(messages).unwrap_or_default()); + if let Some(system) = system { + estimated += estimated_tokens_for_str(system); + } + estimated +} + +pub(crate) fn required_tail_start_for_continuation(history: &[Message]) -> usize { + let Some(last_index) = history.len().checked_sub(1) else { + return 0; + }; + let last_message = &history[last_index]; + + if last_message.role == Role::User + && last_message + .content + .iter() + .any(|block| matches!(block, ContentBlock::ToolResult { .. })) + && last_index > 0 + && history[last_index - 1].role == Role::Assistant + && history[last_index - 1] + .content + .iter() + .any(|block| matches!(block, ContentBlock::ToolUse { .. })) + { + last_index - 1 + } else { + last_index + } +} + +fn tool_name_index(history: &[Message]) -> HashMap { + let mut tool_names = HashMap::new(); + + for message in history { + for block in &message.content { + if let ContentBlock::ToolUse { id, name, .. } = block { + tool_names.insert(id.clone(), name.clone()); + } + } + } + + tool_names +} + +fn estimated_tokens_for_str(text: &str) -> usize { + text.chars().count().div_ceil(4) +} diff --git a/vendor/mentra/src/memory/engine.rs b/vendor/mentra/src/memory/engine.rs new file mode 100644 index 0000000..9f8ca61 --- /dev/null +++ b/vendor/mentra/src/memory/engine.rs @@ -0,0 +1,667 @@ +use std::{ + collections::HashSet, + sync::Arc, + time::{SystemTime, UNIX_EPOCH}, +}; + +use serde::{Deserialize, Serialize}; + +use crate::{ + Message, + provider::ContentBlock, + runtime::{RuntimeError, RuntimeHookEvent, RuntimeHooks, RuntimeStore, TaskItem}, +}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum MemoryRecordKind { + Episode, + Summary, + Fact, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct MemoryRecord { + pub record_id: String, + pub agent_id: String, + pub kind: MemoryRecordKind, + pub content: String, + pub source_revision: u64, + pub created_at: i64, + pub metadata_json: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub source: Option, + #[serde(default, skip_serializing_if = "is_false")] + pub pinned: bool, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub score: Option, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct MemoryCursor { + pub last_ingested_revision: u64, +} + +#[derive(Debug, Clone)] +pub struct SearchRequest { + pub agent_id: String, + pub query: String, + pub limit: usize, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum MemorySearchMode { + #[default] + Automatic, + Tool, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct MemorySearchRequest { + pub agent_id: String, + pub query: String, + pub limit: usize, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub char_budget: Option, + #[serde(default)] + pub mode: MemorySearchMode, + #[serde(default)] + pub filter: MemoryListFilter, +} + +impl From for MemorySearchRequest { + fn from(value: SearchRequest) -> Self { + Self { + agent_id: value.agent_id, + query: value.query, + limit: value.limit, + char_budget: None, + mode: MemorySearchMode::Automatic, + filter: MemoryListFilter::default(), + } + } +} + +#[derive(Debug, Clone, PartialEq)] +pub struct MemoryHit { + pub record_id: String, + pub kind: MemoryRecordKind, + pub content: String, + pub source_revision: u64, + pub created_at: i64, + pub metadata_json: String, + pub source: Option, + pub why_retrieved: Option, + pub score: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum MemoryListSort { + #[default] + Newest, + Oldest, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct MemoryListFilter { + pub kind: Option, + pub pinned: Option, + pub source: Option, + pub created_from: Option, + pub created_to: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct MemoryListCursor { + pub created_at: i64, + pub record_id: String, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct MemoryListRequest { + pub agent_id: String, + pub cursor: Option, + pub limit: usize, + pub filter: MemoryListFilter, + pub sort: MemoryListSort, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct MemoryListPage { + pub records: Vec, + pub next_cursor: Option, +} + +pub const MAX_MEMORY_LIST_PAGE_SIZE: usize = 100; + +#[derive(Debug, Clone)] +pub struct IngestRequest { + pub agent_id: String, + pub source_revision: u64, + pub messages: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct IngestOutcome { + pub stored_records: usize, + pub skipped: bool, +} + +pub trait MemoryStore: Send + Sync { + fn upsert_records(&self, records: &[MemoryRecord]) -> Result<(), RuntimeError>; + fn search_records_with_options( + &self, + request: &MemorySearchRequest, + ) -> Result, RuntimeError>; + fn search_records( + &self, + agent_id: &str, + query: &str, + limit: usize, + ) -> Result, RuntimeError> { + self.search_records_with_options(&MemorySearchRequest { + agent_id: agent_id.to_string(), + query: query.to_string(), + limit, + char_budget: None, + mode: MemorySearchMode::Automatic, + filter: MemoryListFilter::default(), + }) + } + fn list_records(&self, request: &MemoryListRequest) -> Result; + fn get_record( + &self, + agent_id: &str, + record_id: &str, + ) -> Result, RuntimeError>; + fn count_records(&self, agent_id: &str) -> Result; + fn delete_records(&self, record_ids: &[String]) -> Result<(), RuntimeError>; + fn tombstone_records( + &self, + agent_id: &str, + record_ids: &[String], + ) -> Result; + fn load_agent_memory_cursor( + &self, + agent_id: &str, + ) -> Result, RuntimeError>; + fn save_agent_memory_cursor( + &self, + agent_id: &str, + cursor: &MemoryCursor, + ) -> Result<(), RuntimeError>; +} + +#[derive(Clone)] +pub struct MemoryEngine { + store: Arc, + hooks: RuntimeHooks, +} + +impl MemoryEngine { + pub fn new(store: Arc, hooks: RuntimeHooks) -> Self { + Self { store, hooks } + } + + pub async fn search( + &self, + request: impl Into, + ) -> Result, RuntimeError> { + let request = request.into(); + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemorySearchStarted { + agent_id: request.agent_id.clone(), + limit: request.limit, + query_preview: preview_text(&request.query, 120), + }, + ); + let records = match self.store.search_records_with_options(&request) { + Ok(records) => records, + Err(error) => { + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemorySearchFinished { + agent_id: request.agent_id, + success: false, + result_count: 0, + error: Some(error.to_string()), + }, + ); + return Err(error); + } + }; + let mut hits = records + .into_iter() + .map(|record| { + let why_retrieved = build_why_retrieved(&request.query, &record); + MemoryHit { + record_id: record.record_id, + kind: record.kind, + content: record.content, + source_revision: record.source_revision, + created_at: record.created_at, + metadata_json: record.metadata_json, + source: record.source, + why_retrieved, + score: record.score, + } + }) + .collect::>(); + if let Some(char_budget) = request.char_budget { + trim_hits_to_char_budget(&mut hits, char_budget); + } + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemorySearchFinished { + agent_id: request.agent_id, + success: true, + result_count: hits.len(), + error: None, + }, + ); + Ok(hits) + } + + pub fn schedule_ingest(&self, request: IngestRequest) { + let engine = self.clone(); + tokio::spawn(async move { + let _ = engine.ingest(request).await; + }); + } + + pub async fn ingest(&self, request: IngestRequest) -> Result { + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemoryIngestStarted { + agent_id: request.agent_id.clone(), + source_revision: request.source_revision, + }, + ); + + let cursor = match self.store.load_agent_memory_cursor(&request.agent_id) { + Ok(cursor) => cursor.unwrap_or_default(), + Err(error) => { + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemoryIngestFinished { + agent_id: request.agent_id, + source_revision: request.source_revision, + success: false, + stored_records: 0, + error: Some(error.to_string()), + }, + ); + return Err(error); + } + }; + if cursor.last_ingested_revision >= request.source_revision { + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemoryIngestFinished { + agent_id: request.agent_id, + source_revision: request.source_revision, + success: true, + stored_records: 0, + error: None, + }, + ); + return Ok(IngestOutcome { + stored_records: 0, + skipped: true, + }); + } + + let episode = summarize_episode(&request.messages); + if episode.is_empty() { + if let Err(error) = self.store.save_agent_memory_cursor( + &request.agent_id, + &MemoryCursor { + last_ingested_revision: request.source_revision, + }, + ) { + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemoryIngestFinished { + agent_id: request.agent_id, + source_revision: request.source_revision, + success: false, + stored_records: 0, + error: Some(error.to_string()), + }, + ); + return Err(error); + } + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemoryIngestFinished { + agent_id: request.agent_id, + source_revision: request.source_revision, + success: true, + stored_records: 0, + error: None, + }, + ); + return Ok(IngestOutcome { + stored_records: 0, + skipped: false, + }); + } + + let record = MemoryRecord { + record_id: format!("episode:{}:{}", request.agent_id, request.source_revision), + agent_id: request.agent_id.clone(), + kind: MemoryRecordKind::Episode, + content: episode, + source_revision: request.source_revision, + created_at: now_secs(), + metadata_json: "{}".to_string(), + source: Some("auto_ingest".to_string()), + pinned: false, + score: None, + }; + if let Err(error) = self.store.upsert_records(&[record]) { + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemoryIngestFinished { + agent_id: request.agent_id, + source_revision: request.source_revision, + success: false, + stored_records: 0, + error: Some(error.to_string()), + }, + ); + return Err(error); + } + if let Err(error) = self.store.save_agent_memory_cursor( + &request.agent_id, + &MemoryCursor { + last_ingested_revision: request.source_revision, + }, + ) { + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemoryIngestFinished { + agent_id: request.agent_id, + source_revision: request.source_revision, + success: false, + stored_records: 0, + error: Some(error.to_string()), + }, + ); + return Err(error); + } + let _ = self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::MemoryIngestFinished { + agent_id: request.agent_id, + source_revision: request.source_revision, + success: true, + stored_records: 1, + error: None, + }, + ); + Ok(IngestOutcome { + stored_records: 1, + skipped: false, + }) + } + + pub fn store_compaction_summary( + &self, + agent_id: &str, + source_revision: u64, + summary: &str, + ) -> Result<(), RuntimeError> { + let record = MemoryRecord { + record_id: format!("summary:{agent_id}:{source_revision}"), + agent_id: agent_id.to_string(), + kind: MemoryRecordKind::Summary, + content: summary.to_string(), + source_revision, + created_at: now_secs(), + metadata_json: "{}".to_string(), + source: Some("auto_compaction".to_string()), + pinned: false, + score: None, + }; + self.store.upsert_records(&[record]) + } + + pub fn pin( + &self, + agent_id: &str, + source_revision: u64, + content: &str, + ) -> Result { + let record = MemoryRecord { + record_id: format!("fact:{agent_id}:manual:{}", now_nanos()), + agent_id: agent_id.to_string(), + kind: MemoryRecordKind::Fact, + content: content.trim().to_string(), + source_revision, + created_at: now_secs(), + metadata_json: r#"{"origin":"manual_pin"}"#.to_string(), + source: Some("manual_pin".to_string()), + pinned: true, + score: None, + }; + self.store.upsert_records(std::slice::from_ref(&record))?; + Ok(record) + } + + pub fn forget(&self, agent_id: &str, record_id: &str) -> Result { + self.store + .tombstone_records(agent_id, &[record_id.to_string()]) + .map(|count| count > 0) + } +} + +pub(crate) fn build_search_query(history: &[Message], tasks: &[TaskItem]) -> String { + let mut parts = history.iter().rev().take(6).collect::>(); + parts.reverse(); + + let mut query = parts + .into_iter() + .flat_map(message_to_lines) + .collect::>() + .join("\n"); + + let unfinished = tasks + .iter() + .filter(|task| !matches!(task.status, crate::runtime::TaskStatus::Completed)) + .map(|task| { + let description = task.description.trim(); + if description.is_empty() { + task.subject.clone() + } else { + format!("{}: {}", task.subject, description) + } + }) + .collect::>(); + if !unfinished.is_empty() { + if !query.is_empty() { + query.push('\n'); + } + query.push_str("Tasks:\n"); + query.push_str(&unfinished.join("\n")); + } + query +} + +pub(crate) fn recalled_memory_message(hits: &[MemoryHit], char_limit: usize) -> Option { + let mut seen = HashSet::new(); + let mut entries = Vec::new(); + let mut used = 0usize; + + for hit in hits { + if !seen.insert(hit.record_id.clone()) { + continue; + } + let line = format!( + "[{} rev={}{}{}] {}", + kind_label(hit.kind), + hit.source_revision, + hit.source + .as_deref() + .map(|source| format!(" source={source}")) + .unwrap_or_default(), + hit.why_retrieved + .as_deref() + .map(|why| format!(" why={why}")) + .unwrap_or_default(), + hit.content.trim() + ); + if line.trim().is_empty() { + continue; + } + if !entries.is_empty() && used + line.len() + 1 > char_limit { + break; + } + used += if entries.is_empty() { + line.len() + } else { + line.len() + 1 + }; + entries.push(line); + } + + if entries.is_empty() { + return None; + } + + Some(Message::user(ContentBlock::text(format!( + "\n{}\n", + entries.join("\n") + )))) +} + +fn summarize_episode(messages: &[Message]) -> String { + let mut lines = Vec::new(); + for message in messages { + let label = match message.role { + crate::Role::User => "user", + crate::Role::Assistant => "assistant", + crate::Role::Unknown(_) => "unknown", + }; + for line in message_to_lines(message) { + lines.push(format!("{label}: {line}")); + } + } + lines.join("\n") +} + +fn message_to_lines(message: &Message) -> Vec { + message + .content + .iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text.trim().to_string()), + ContentBlock::ToolUse { name, input, .. } => Some(format!("tool use {name} {input}")), + ContentBlock::ToolResult { content, .. } => Some(format!("tool result {content}")), + ContentBlock::Image { .. } + | ContentBlock::Thinking { .. } + | ContentBlock::HostedToolSearch { .. } + | ContentBlock::HostedWebSearch { .. } + | ContentBlock::ImageGeneration { .. } => None, + }) + .filter(|text| !text.is_empty()) + .collect() +} + +fn kind_label(kind: MemoryRecordKind) -> &'static str { + match kind { + MemoryRecordKind::Episode => "episode", + MemoryRecordKind::Summary => "summary", + MemoryRecordKind::Fact => "fact", + } +} + +fn preview_text(text: &str, limit: usize) -> String { + truncate_to_char_boundary(text.trim(), limit).to_string() +} + +fn truncate_to_char_boundary(input: &str, max_chars: usize) -> &str { + if input.chars().count() <= max_chars { + return input; + } + + let mut end = input.len(); + for (index, _) in input.char_indices().take(max_chars + 1) { + end = index; + } + &input[..end] +} + +fn now_secs() -> i64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64 +} + +fn now_nanos() -> u128 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() +} + +fn trim_hits_to_char_budget(hits: &mut Vec, char_budget: usize) { + if char_budget == 0 { + hits.clear(); + return; + } + + let mut kept = Vec::with_capacity(hits.len()); + let mut used = 0usize; + for hit in hits.drain(..) { + let line_len = hit.content.len() + + hit.source.as_deref().map_or(0, str::len) + + hit.why_retrieved.as_deref().map_or(0, str::len); + if !kept.is_empty() && used + line_len > char_budget { + break; + } + used += line_len; + kept.push(hit); + } + *hits = kept; +} + +fn build_why_retrieved(query: &str, record: &MemoryRecord) -> Option { + let mut reasons = Vec::new(); + let matched = query + .split(|ch: char| !ch.is_alphanumeric()) + .filter(|token| !token.is_empty()) + .filter(|token| { + record + .content + .to_lowercase() + .contains(&token.to_lowercase()) + }) + .take(2) + .map(ToString::to_string) + .collect::>(); + if !matched.is_empty() { + reasons.push(format!("matched {}", matched.join(","))); + } + match record.kind { + MemoryRecordKind::Fact => reasons.push("fact".to_string()), + MemoryRecordKind::Summary => reasons.push("summary".to_string()), + MemoryRecordKind::Episode => {} + } + if record.pinned || record.source.as_deref() == Some("manual_pin") { + reasons.push("manual".to_string()); + } + if reasons.is_empty() { + None + } else { + Some(reasons.join("; ")) + } +} + +fn is_false(value: &bool) -> bool { + !*value +} diff --git a/vendor/mentra/src/memory/hybrid_store.rs b/vendor/mentra/src/memory/hybrid_store.rs new file mode 100644 index 0000000..ab40ae7 --- /dev/null +++ b/vendor/mentra/src/memory/hybrid_store.rs @@ -0,0 +1,861 @@ +use std::{ + path::{Path, PathBuf}, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use rusqlite::{Connection, OptionalExtension, TransactionBehavior, params}; + +use crate::{ + memory::{ + MemoryCursor, MemoryListCursor, MemoryListPage, MemoryListRequest, MemoryListSort, + MemoryRecord, MemoryRecordKind, MemorySearchRequest, MemoryStore, + }, + runtime::RuntimeError, +}; + +#[derive(Clone)] +/// SQLite-backed hybrid memory store with provenance, pinning, and tombstoning support. +pub struct SqliteHybridMemoryStore { + path: PathBuf, +} + +impl SqliteHybridMemoryStore { + pub fn new(path: impl Into) -> Self { + Self { path: path.into() } + } + + pub fn path(&self) -> &Path { + self.path.as_path() + } + + fn open(&self) -> Result { + if let Some(parent) = self.path.parent() { + std::fs::create_dir_all(parent) + .map_err(|error| RuntimeError::Store(error.to_string()))?; + } + let conn = Connection::open(&self.path).map_err(sqlite_error)?; + conn.busy_timeout(Duration::from_secs(5)) + .map_err(sqlite_error)?; + conn.pragma_update(None, "journal_mode", "WAL") + .map_err(sqlite_error)?; + self.ensure_schema(&conn)?; + Ok(conn) + } + + fn ensure_schema(&self, conn: &Connection) -> Result<(), RuntimeError> { + conn.execute_batch( + r#" + CREATE TABLE IF NOT EXISTS memory_records ( + record_id TEXT PRIMARY KEY, + agent_id TEXT NOT NULL, + kind TEXT NOT NULL, + content TEXT NOT NULL, + source_revision INTEGER NOT NULL, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + metadata_json TEXT NOT NULL, + source_json TEXT, + pinned INTEGER NOT NULL DEFAULT 0, + tombstoned_at INTEGER + ); + CREATE INDEX IF NOT EXISTS idx_memory_records_agent_created + ON memory_records (agent_id, created_at DESC); + CREATE VIRTUAL TABLE IF NOT EXISTS memory_records_fts USING fts5( + record_id UNINDEXED, + agent_id UNINDEXED, + content + ); + CREATE TABLE IF NOT EXISTS memory_cursor ( + agent_id TEXT PRIMARY KEY, + cursor_json TEXT NOT NULL, + updated_at INTEGER NOT NULL + ); + "#, + ) + .map_err(sqlite_error) + } + + fn search_records_raw( + &self, + request: &MemorySearchRequest, + ) -> Result, RuntimeError> { + if request.query.trim().is_empty() || request.limit == 0 { + return Ok(Vec::new()); + } + let Some(query) = fts_query(&request.query) else { + return Ok(Vec::new()); + }; + + let conn = self.open()?; + let kind = request.filter.kind.map(kind_name); + let source = encode_source(request.filter.source.as_deref())?; + let mut stmt = conn + .prepare( + r#" + SELECT + record.record_id, + record.agent_id, + record.kind, + record.content, + record.source_revision, + record.created_at, + record.metadata_json, + record.source_json, + record.pinned, + bm25(memory_records_fts) AS rank + FROM memory_records_fts + JOIN memory_records AS record ON record.record_id = memory_records_fts.record_id + WHERE memory_records_fts.agent_id = ?1 + AND memory_records_fts.content MATCH ?2 + AND record.tombstoned_at IS NULL + AND (?3 IS NULL OR record.kind = ?3) + AND (?4 IS NULL OR record.pinned = ?4) + AND (?5 IS NULL OR record.source_json = ?5) + AND (?6 IS NULL OR record.created_at >= ?6) + AND (?7 IS NULL OR record.created_at <= ?7) + LIMIT ?8 + "#, + ) + .map_err(sqlite_error)?; + + let candidate_limit = request.limit.saturating_mul(5).clamp(10, 500) as i64; + let mut records = stmt + .query_map( + params![ + request.agent_id, + query, + kind, + request.filter.pinned.map(i64::from), + source, + request.filter.created_from, + request.filter.created_to, + candidate_limit, + ], + |row| { + let kind = row.get::<_, String>(2)?; + let source_json = row.get::<_, Option>(7)?; + let pinned = row.get::<_, i64>(8)? != 0; + let raw_rank = row.get::<_, Option>(9)?.unwrap_or(0.0); + let created_at = row.get::<_, i64>(5)?; + let score = rank_score(parse_memory_kind(&kind), pinned, created_at, raw_rank); + Ok(MemoryRecord { + record_id: row.get(0)?, + agent_id: row.get(1)?, + kind: parse_memory_kind(&kind), + content: row.get(3)?, + source_revision: row.get::<_, i64>(4)? as u64, + created_at, + metadata_json: row.get(6)?, + source: decode_source(source_json), + pinned, + score: Some(score), + }) + }, + ) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error)?; + + records.sort_by(|left, right| { + right + .score + .partial_cmp(&left.score) + .unwrap_or(std::cmp::Ordering::Equal) + .then_with(|| right.created_at.cmp(&left.created_at)) + }); + records.truncate(request.limit); + Ok(records) + } +} + +impl MemoryStore for SqliteHybridMemoryStore { + fn upsert_records(&self, records: &[MemoryRecord]) -> Result<(), RuntimeError> { + if records.is_empty() { + return Ok(()); + } + + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + let now = now_secs(); + + for record in records { + tx.execute( + r#" + INSERT INTO memory_records ( + record_id, agent_id, kind, content, source_revision, created_at, updated_at, + metadata_json, source_json, pinned, tombstoned_at + ) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, NULL) + ON CONFLICT(record_id) DO UPDATE SET + agent_id = excluded.agent_id, + kind = excluded.kind, + content = excluded.content, + source_revision = excluded.source_revision, + created_at = excluded.created_at, + updated_at = excluded.updated_at, + metadata_json = excluded.metadata_json, + source_json = excluded.source_json, + pinned = excluded.pinned, + tombstoned_at = NULL + "#, + params![ + record.record_id, + record.agent_id, + kind_name(record.kind), + record.content, + record.source_revision as i64, + record.created_at, + now, + record.metadata_json, + encode_source(record.source.as_deref())?, + if record.pinned { 1 } else { 0 }, + ], + ) + .map_err(sqlite_error)?; + tx.execute( + "DELETE FROM memory_records_fts WHERE record_id = ?1", + params![record.record_id], + ) + .map_err(sqlite_error)?; + tx.execute( + "INSERT INTO memory_records_fts (record_id, agent_id, content) VALUES (?1, ?2, ?3)", + params![record.record_id, record.agent_id, record.content], + ) + .map_err(sqlite_error)?; + } + + tx.commit().map_err(sqlite_error) + } + + fn search_records_with_options( + &self, + request: &MemorySearchRequest, + ) -> Result, RuntimeError> { + self.search_records_raw(request) + } + + fn list_records(&self, request: &MemoryListRequest) -> Result { + let limit = request.limit.min(crate::memory::MAX_MEMORY_LIST_PAGE_SIZE); + if limit == 0 { + return Ok(MemoryListPage { + records: Vec::new(), + next_cursor: None, + }); + } + let conn = self.open()?; + let kind = request.filter.kind.map(kind_name); + let source = encode_source(request.filter.source.as_deref())?; + let cursor_time = request.cursor.as_ref().map(|cursor| cursor.created_at); + let cursor_id = request + .cursor + .as_ref() + .map(|cursor| cursor.record_id.as_str()); + let direction = match request.sort { + MemoryListSort::Newest => "<", + MemoryListSort::Oldest => ">", + }; + let order = match request.sort { + MemoryListSort::Newest => "DESC", + MemoryListSort::Oldest => "ASC", + }; + let sql = format!( + r#" + SELECT record_id, agent_id, kind, content, source_revision, created_at, + metadata_json, source_json, pinned + FROM memory_records + WHERE agent_id = ?1 AND tombstoned_at IS NULL + AND (?2 IS NULL OR kind = ?2) + AND (?3 IS NULL OR pinned = ?3) + AND (?4 IS NULL OR source_json = ?4) + AND (?5 IS NULL OR created_at >= ?5) + AND (?6 IS NULL OR created_at <= ?6) + AND (?7 IS NULL OR created_at {direction} ?7 + OR (created_at = ?7 AND record_id {direction} ?8)) + ORDER BY created_at {order}, record_id {order} + LIMIT ?9 + "# + ); + let mut stmt = conn.prepare(&sql).map_err(sqlite_error)?; + let mut records = stmt + .query_map( + params![ + request.agent_id, + kind, + request.filter.pinned.map(i64::from), + source, + request.filter.created_from, + request.filter.created_to, + cursor_time, + cursor_id, + limit.saturating_add(1) as i64, + ], + memory_record_from_row, + ) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error)?; + let has_more = records.len() > limit; + records.truncate(limit); + let next_cursor = if has_more { + records.last().map(|record| MemoryListCursor { + created_at: record.created_at, + record_id: record.record_id.clone(), + }) + } else { + None + }; + Ok(MemoryListPage { + records, + next_cursor, + }) + } + + fn get_record( + &self, + agent_id: &str, + record_id: &str, + ) -> Result, RuntimeError> { + let conn = self.open()?; + conn.query_row( + r#" + SELECT record_id, agent_id, kind, content, source_revision, created_at, + metadata_json, source_json, pinned + FROM memory_records + WHERE agent_id = ?1 AND record_id = ?2 AND tombstoned_at IS NULL + "#, + params![agent_id, record_id], + memory_record_from_row, + ) + .optional() + .map_err(sqlite_error) + } + + fn count_records(&self, agent_id: &str) -> Result { + let conn = self.open()?; + conn.query_row( + "SELECT COUNT(*) FROM memory_records WHERE agent_id = ?1 AND tombstoned_at IS NULL", + params![agent_id], + |row| row.get::<_, i64>(0), + ) + .map(|count| count as usize) + .map_err(sqlite_error) + } + + fn delete_records(&self, record_ids: &[String]) -> Result<(), RuntimeError> { + if record_ids.is_empty() { + return Ok(()); + } + + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + for record_id in record_ids { + tx.execute( + "DELETE FROM memory_records_fts WHERE record_id = ?1", + params![record_id], + ) + .map_err(sqlite_error)?; + tx.execute( + "DELETE FROM memory_records WHERE record_id = ?1", + params![record_id], + ) + .map_err(sqlite_error)?; + } + tx.commit().map_err(sqlite_error) + } + + fn tombstone_records( + &self, + agent_id: &str, + record_ids: &[String], + ) -> Result { + if record_ids.is_empty() { + return Ok(0); + } + + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + let mut affected = 0usize; + let now = now_secs(); + + for record_id in record_ids { + let updated = tx + .execute( + r#" + UPDATE memory_records + SET tombstoned_at = ?3, updated_at = ?3 + WHERE record_id = ?1 AND agent_id = ?2 AND tombstoned_at IS NULL + "#, + params![record_id, agent_id, now], + ) + .map_err(sqlite_error)?; + if updated > 0 { + affected += updated; + tx.execute( + "DELETE FROM memory_records_fts WHERE record_id = ?1", + params![record_id], + ) + .map_err(sqlite_error)?; + } + } + + tx.commit().map_err(sqlite_error)?; + Ok(affected) + } + + fn load_agent_memory_cursor( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + let conn = self.open()?; + conn.query_row( + "SELECT cursor_json FROM memory_cursor WHERE agent_id = ?1", + params![agent_id], + |row| row.get::<_, String>(0), + ) + .optional() + .map_err(sqlite_error)? + .map(|json| from_json(&json)) + .transpose() + } + + fn save_agent_memory_cursor( + &self, + agent_id: &str, + cursor: &MemoryCursor, + ) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + r#" + INSERT INTO memory_cursor (agent_id, cursor_json, updated_at) + VALUES (?1, ?2, ?3) + ON CONFLICT(agent_id) DO UPDATE SET + cursor_json = excluded.cursor_json, + updated_at = excluded.updated_at + "#, + params![agent_id, to_json(cursor)?, now_secs()], + ) + .map_err(sqlite_error)?; + Ok(()) + } +} + +fn memory_record_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result { + let kind = row.get::<_, String>(2)?; + Ok(MemoryRecord { + record_id: row.get(0)?, + agent_id: row.get(1)?, + kind: parse_memory_kind(&kind), + content: row.get(3)?, + source_revision: row.get::<_, i64>(4)? as u64, + created_at: row.get(5)?, + metadata_json: row.get(6)?, + source: decode_source(row.get(7)?), + pinned: row.get::<_, i64>(8)? != 0, + score: None, + }) +} + +fn parse_memory_kind(kind: &str) -> MemoryRecordKind { + match kind { + "summary" => MemoryRecordKind::Summary, + "fact" => MemoryRecordKind::Fact, + _ => MemoryRecordKind::Episode, + } +} + +fn kind_name(kind: MemoryRecordKind) -> &'static str { + match kind { + MemoryRecordKind::Episode => "episode", + MemoryRecordKind::Summary => "summary", + MemoryRecordKind::Fact => "fact", + } +} + +fn rank_score(kind: MemoryRecordKind, pinned: bool, created_at: i64, raw_rank: f64) -> f64 { + let kind_bonus = match kind { + MemoryRecordKind::Fact => 3.0, + MemoryRecordKind::Summary => 1.5, + MemoryRecordKind::Episode => 0.0, + }; + let manual_bonus = if pinned { 2.0 } else { 0.0 }; + let age_hours = ((now_secs() - created_at).max(0) as f64) / 3600.0; + let recency_bonus = 0.5 / (1.0 + age_hours / 24.0); + let text_bonus = 8.0 / (1.0 + raw_rank.abs()); + text_bonus + kind_bonus + manual_bonus + recency_bonus +} + +fn encode_source(source: Option<&str>) -> Result, RuntimeError> { + source + .map(|value| { + serde_json::to_string(value).map_err(|error| RuntimeError::Store(error.to_string())) + }) + .transpose() +} + +fn decode_source(source_json: Option) -> Option { + source_json.and_then(|json| serde_json::from_str::(&json).ok()) +} + +fn to_json(value: &T) -> Result { + serde_json::to_string(value).map_err(|error| RuntimeError::Store(error.to_string())) +} + +fn from_json(value: &str) -> Result { + serde_json::from_str(value).map_err(|error| RuntimeError::Store(error.to_string())) +} + +fn sqlite_error(error: rusqlite::Error) -> RuntimeError { + RuntimeError::Store(error.to_string()) +} + +fn now_secs() -> i64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64 +} + +fn fts_query(query: &str) -> Option { + let tokens = query + .split(|ch: char| !ch.is_alphanumeric()) + .filter(|token| !token.is_empty()) + .map(|token| format!("\"{token}\"")) + .collect::>(); + + if tokens.is_empty() { + None + } else { + Some(tokens.join(" OR ")) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::memory::{MemoryListFilter, MemorySearchMode, MemoryStore}; + + #[test] + fn pinned_manual_facts_outrank_episodes() { + let store = SqliteHybridMemoryStore::new( + std::env::temp_dir().join(format!("mentra-hybrid-memory-{}.sqlite", now_secs())), + ); + store + .upsert_records(&[ + MemoryRecord { + record_id: "episode:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Episode, + content: "shared phrase alpha".to_string(), + source_revision: 1, + created_at: now_secs(), + metadata_json: "{}".to_string(), + source: Some("auto_ingest".to_string()), + pinned: false, + score: None, + }, + MemoryRecord { + record_id: "fact:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Fact, + content: "shared phrase alpha".to_string(), + source_revision: 2, + created_at: now_secs(), + metadata_json: "{}".to_string(), + source: Some("manual_pin".to_string()), + pinned: true, + score: None, + }, + ]) + .expect("seed records"); + + let records = store + .search_records_with_options(&MemorySearchRequest { + agent_id: "agent-1".to_string(), + query: "shared alpha".to_string(), + limit: 2, + char_budget: None, + mode: MemorySearchMode::Tool, + filter: MemoryListFilter { + kind: Some(MemoryRecordKind::Fact), + ..MemoryListFilter::default() + }, + }) + .expect("search"); + assert_eq!(records[0].record_id, "fact:1"); + assert!( + records + .iter() + .all(|record| record.kind == MemoryRecordKind::Fact) + ); + } + + #[test] + fn tombstoned_records_are_excluded_from_reads() { + let store = SqliteHybridMemoryStore::new( + std::env::temp_dir().join(format!("mentra-hybrid-tombstone-{}.sqlite", now_secs())), + ); + store + .upsert_records(&[MemoryRecord { + record_id: "fact:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Fact, + content: "preferred editor is vim".to_string(), + source_revision: 1, + created_at: now_secs(), + metadata_json: "{}".to_string(), + source: Some("manual_pin".to_string()), + pinned: true, + score: None, + }]) + .expect("seed records"); + assert_eq!( + store + .tombstone_records("agent-1", &["fact:1".to_string()]) + .expect("tombstone"), + 1 + ); + + let records = store.search_records("agent-1", "vim", 5).expect("search"); + assert!(records.is_empty()); + } + + #[test] + fn punctuation_heavy_queries_still_return_results() { + let store = SqliteHybridMemoryStore::new( + std::env::temp_dir().join(format!("mentra-hybrid-punct-{}.sqlite", now_secs())), + ); + store + .upsert_records(&[MemoryRecord { + record_id: "episode:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Episode, + content: "shared phrase alpha".to_string(), + source_revision: 1, + created_at: now_secs(), + metadata_json: "{}".to_string(), + source: Some("auto_ingest".to_string()), + pinned: false, + score: None, + }]) + .expect("seed records"); + + let records = store + .search_records("agent-1", "(shared) alpha!!!", 5) + .expect("search"); + assert_eq!(records.len(), 1); + } + + #[test] + fn compatibility_search_wrapper_matches_options_search() { + let store = SqliteHybridMemoryStore::new( + std::env::temp_dir().join(format!("mentra-hybrid-compat-{}.sqlite", now_secs())), + ); + store + .upsert_records(&[MemoryRecord { + record_id: "episode:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Episode, + content: "shared phrase alpha".to_string(), + source_revision: 1, + created_at: now_secs(), + metadata_json: "{}".to_string(), + source: Some("auto_ingest".to_string()), + pinned: false, + score: None, + }]) + .expect("seed records"); + + let compat = store + .search_records("agent-1", "shared alpha", 5) + .expect("compat search"); + let explicit = store + .search_records_with_options(&MemorySearchRequest { + agent_id: "agent-1".to_string(), + query: "shared alpha".to_string(), + limit: 5, + char_budget: None, + mode: MemorySearchMode::Automatic, + filter: MemoryListFilter::default(), + }) + .expect("explicit search"); + assert_eq!(compat.len(), explicit.len()); + for (compat_record, explicit_record) in compat.iter().zip(explicit.iter()) { + assert_eq!(compat_record.record_id, explicit_record.record_id); + assert_eq!(compat_record.agent_id, explicit_record.agent_id); + assert_eq!(compat_record.kind, explicit_record.kind); + assert_eq!(compat_record.content, explicit_record.content); + assert_eq!( + compat_record.source_revision, + explicit_record.source_revision + ); + assert_eq!(compat_record.created_at, explicit_record.created_at); + assert_eq!(compat_record.metadata_json, explicit_record.metadata_json); + assert_eq!(compat_record.source, explicit_record.source); + assert_eq!(compat_record.pinned, explicit_record.pinned); + + let compat_score = compat_record.score.expect("compat score"); + let explicit_score = explicit_record.score.expect("explicit score"); + assert!( + (compat_score - explicit_score).abs() < 1e-5, + "expected comparable ranking scores, got {compat_score} vs {explicit_score}" + ); + } + } + + #[test] + fn stable_pages_filters_and_tombstones_survive_restart() { + let path = std::env::temp_dir().join(format!( + "mentra-hybrid-list-{}-{}.sqlite", + std::process::id(), + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos() + )); + let store = SqliteHybridMemoryStore::new(&path); + let make = |id: &str, agent: &str, kind, pinned, source: &str, created_at| MemoryRecord { + record_id: id.to_owned(), + agent_id: agent.to_owned(), + kind, + content: format!("content {id}"), + source_revision: 1, + created_at, + metadata_json: "{}".to_owned(), + source: Some(source.to_owned()), + pinned, + score: None, + }; + store + .upsert_records(&[ + make( + "fact:c", + "agent-1", + MemoryRecordKind::Fact, + true, + "manual", + 20, + ), + make( + "fact:b", + "agent-1", + MemoryRecordKind::Fact, + true, + "manual", + 20, + ), + make( + "episode:a", + "agent-1", + MemoryRecordKind::Episode, + false, + "auto", + 10, + ), + make( + "fact:other", + "agent-2", + MemoryRecordKind::Fact, + true, + "manual", + 30, + ), + ]) + .unwrap(); + let bulk = (0..105) + .map(|index| { + make( + &format!("episode:bulk:{index:03}"), + "agent-bulk", + MemoryRecordKind::Episode, + false, + "auto", + index, + ) + }) + .collect::>(); + store.upsert_records(&bulk).unwrap(); + let bounded = store + .list_records(&MemoryListRequest { + agent_id: "agent-bulk".to_owned(), + cursor: None, + limit: usize::MAX, + filter: MemoryListFilter::default(), + sort: MemoryListSort::Newest, + }) + .unwrap(); + assert_eq!(bounded.records.len(), 100); + assert!(bounded.next_cursor.is_some()); + let request = MemoryListRequest { + agent_id: "agent-1".to_owned(), + cursor: None, + limit: 1, + filter: MemoryListFilter { + kind: Some(MemoryRecordKind::Fact), + pinned: Some(true), + source: Some("manual".to_owned()), + ..MemoryListFilter::default() + }, + sort: MemoryListSort::Newest, + }; + let first = store.list_records(&request).unwrap(); + assert_eq!(first.records[0].record_id, "fact:c"); + store + .upsert_records(&[make( + "fact:d", + "agent-1", + MemoryRecordKind::Fact, + true, + "manual", + 30, + )]) + .unwrap(); + let second = store + .list_records(&MemoryListRequest { + cursor: first.next_cursor, + ..request + }) + .unwrap(); + assert_eq!(second.records[0].record_id, "fact:b"); + assert_eq!(store.count_records("agent-1").unwrap(), 4); + assert!(store.get_record("agent-2", "fact:b").unwrap().is_none()); + assert_eq!( + store + .tombstone_records("agent-1", &["fact:b".to_owned()]) + .unwrap(), + 1 + ); + drop(store); + let reopened = SqliteHybridMemoryStore::new(&path); + assert!(reopened.get_record("agent-1", "fact:b").unwrap().is_none()); + assert_eq!(reopened.count_records("agent-1").unwrap(), 3); + let after_restart = reopened + .list_records(&MemoryListRequest { + agent_id: "agent-1".to_owned(), + cursor: None, + limit: 100, + filter: MemoryListFilter::default(), + sort: MemoryListSort::Newest, + }) + .unwrap(); + assert!( + after_restart + .records + .iter() + .all(|record| record.record_id != "fact:b") + ); + assert!( + reopened + .search_records("agent-1", "fact b", 100) + .unwrap() + .iter() + .all(|record| record.record_id != "fact:b") + ); + let _ = std::fs::remove_file(path); + } +} diff --git a/vendor/mentra/src/memory/journal.rs b/vendor/mentra/src/memory/journal.rs new file mode 100644 index 0000000..83163ec --- /dev/null +++ b/vendor/mentra/src/memory/journal.rs @@ -0,0 +1,10 @@ +mod ops; +mod recovery; +mod snapshot; +mod state; +mod store; +#[cfg(test)] +mod tests; + +pub(crate) use ops::{AgentMemory, CompactionOutcome}; +pub(crate) use state::{AgentMemoryState, PendingTurnState}; diff --git a/vendor/mentra/src/memory/journal/ops.rs b/vendor/mentra/src/memory/journal/ops.rs new file mode 100644 index 0000000..c707859 --- /dev/null +++ b/vendor/mentra/src/memory/journal/ops.rs @@ -0,0 +1,232 @@ +use std::{collections::BTreeMap, path::PathBuf, sync::Arc}; + +use crate::{ + Message, + error::RuntimeError, + runtime::RuntimeStore, + transcript::{AgentTranscript, EntryId, TranscriptItem, transcript_item_from_message}, +}; + +use super::{ + recovery::RecoveryOutcome, + snapshot::AgentSnapshotMemoryView, + state::{AgentMemoryState, PendingTurnState, RunMemoryState}, + store::AgentMemoryStore, +}; + +#[derive(Debug, Clone)] +pub(crate) struct CompactionOutcome { + pub transcript_path: PathBuf, + pub transcript: AgentTranscript, +} + +pub(crate) struct AgentMemory { + agent_id: String, + store: Arc, + state: AgentMemoryState, + history_cache: Vec, +} + +impl AgentMemory { + pub fn new( + agent_id: impl Into, + store: Arc, + state: AgentMemoryState, + ) -> Self { + let history_cache = state.transcript.to_messages(); + Self { + agent_id: agent_id.into(), + store, + state, + history_cache, + } + } + + pub fn begin_run(&mut self, run_id: String, user_message: Message) -> Result<(), RuntimeError> { + self.state.run = Some(RunMemoryState { + run_id, + baseline_transcript: self.state.transcript.clone(), + assistant_committed: false, + }); + self.state.pending_turn = None; + self.state.resumable_user_message = Some(user_message.clone()); + self.state + .transcript + .push(transcript_item_from_message(user_message)); + self.sync_history_cache(); + self.persist() + } + + pub fn append_message(&mut self, message: Message) -> Result<(), RuntimeError> { + self.append_transcript_item(transcript_item_from_message(message)) + } + + /// Additive counterpart to [`Self::append_message`] that also attaches + /// opaque per-call host metadata (keyed by `tool_use_id`) to the + /// resulting transcript item, so it survives persistence and replay + /// without mentra interpreting it (ADR-0001 §4). + pub fn append_message_with_details( + &mut self, + message: Message, + details: BTreeMap, + ) -> Result<(), RuntimeError> { + self.append_transcript_item(transcript_item_from_message(message).with_details(details)) + } + + pub fn append_transcript_item(&mut self, item: TranscriptItem) -> Result<(), RuntimeError> { + self.state.transcript.push(item); + self.sync_history_cache(); + self.persist() + } + + pub fn update_pending_turn(&mut self, pending: PendingTurnState) -> Result<(), RuntimeError> { + self.state.pending_turn = Some(pending); + self.persist() + } + + pub fn clear_pending_turn(&mut self) -> Result<(), RuntimeError> { + self.state.pending_turn = None; + self.persist() + } + + pub fn commit_assistant_message(&mut self, message: Message) -> Result<(), RuntimeError> { + self.state + .transcript + .push(transcript_item_from_message(message)); + self.sync_history_cache(); + self.state.pending_turn = None; + if let Some(run) = &mut self.state.run { + run.assistant_committed = true; + } + self.persist() + } + + #[cfg(test)] + pub fn compact(&mut self, outcome: CompactionOutcome) -> Result<(), RuntimeError> { + self.state.transcript = outcome.transcript; + self.sync_history_cache(); + let _ = outcome.transcript_path; + self.persist() + } + + pub fn rollback_failed_run(&mut self) -> Result<(), RuntimeError> { + if let Some(run) = self.state.run.take() { + self.state.transcript = run.baseline_transcript; + } + self.sync_history_cache(); + self.state.pending_turn = None; + self.persist() + } + + pub fn finish_run(&mut self) -> Result<(), RuntimeError> { + self.state.pending_turn = None; + self.state.run = None; + self.state.resumable_user_message = None; + self.persist() + } + + pub fn recover(&mut self) -> Result { + let Some(run) = self.state.run.take() else { + return Ok(RecoveryOutcome::default()); + }; + + let had_pending_turn = self.state.pending_turn.take().is_some(); + if had_pending_turn || !run.assistant_committed { + self.state.transcript = run.baseline_transcript; + self.sync_history_cache(); + } else { + self.state.resumable_user_message = None; + } + + self.persist()?; + Ok(RecoveryOutcome { + interrupted: true, + interrupted_run_id: Some(run.run_id), + }) + } + + /// Moves the transcript leaf back to `id` and re-derives history. + pub fn branch_from(&mut self, id: &EntryId) -> Result { + let moved = self + .state + .transcript + .branch_from(id) + .map_err(RuntimeError::Branch)?; + self.sync_history_cache(); + self.persist()?; + Ok(moved) + } + + pub fn transcript(&self) -> &AgentTranscript { + &self.state.transcript + } + + pub fn history(&self) -> &[Message] { + &self.history_cache + } + + pub fn revision(&self) -> u64 { + self.state.revision + } + + pub fn last_message(&self) -> Option<&Message> { + self.history_cache.last() + } + + pub fn resumable_user_message(&self) -> Option<&Message> { + self.state.resumable_user_message.as_ref() + } + + pub fn snapshot_view(&self) -> AgentSnapshotMemoryView { + AgentSnapshotMemoryView::from(&self.state) + } + + pub fn state(&self) -> &AgentMemoryState { + &self.state + } + + pub fn current_run_delta(&self) -> Option> { + let run = self.state.run.as_ref()?; + let start = run.baseline_transcript.len(); + if start >= self.state.transcript.len() { + return Some(self.history_cache.clone()); + } + Some(self.state.transcript.projected_messages_from(start)) + } + + pub fn try_apply_compaction( + &mut self, + base_revision: u64, + outcome: CompactionOutcome, + ) -> Result { + if self.state.revision != base_revision { + return Ok(false); + } + self.state.transcript = outcome.transcript; + self.sync_history_cache(); + let _ = outcome.transcript_path; + self.persist()?; + Ok(true) + } + + fn sync_history_cache(&mut self) { + self.history_cache = self.state.transcript.to_messages(); + } + + fn persist(&mut self) -> Result<(), RuntimeError> { + self.state.revision = self.state.revision.saturating_add(1); + self.store.save_memory(&self.agent_id, &self.state) + } +} + +impl PendingTurnState { + pub fn new( + current_text: String, + pending_tool_uses: Vec, + ) -> Self { + Self { + current_text, + pending_tool_uses, + } + } +} diff --git a/vendor/mentra/src/memory/journal/recovery.rs b/vendor/mentra/src/memory/journal/recovery.rs new file mode 100644 index 0000000..62fb21d --- /dev/null +++ b/vendor/mentra/src/memory/journal/recovery.rs @@ -0,0 +1,5 @@ +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct RecoveryOutcome { + pub interrupted: bool, + pub interrupted_run_id: Option, +} diff --git a/vendor/mentra/src/memory/journal/snapshot.rs b/vendor/mentra/src/memory/journal/snapshot.rs new file mode 100644 index 0000000..3624685 --- /dev/null +++ b/vendor/mentra/src/memory/journal/snapshot.rs @@ -0,0 +1,28 @@ +use crate::agent::PendingToolUseSummary; + +use super::state::AgentMemoryState; + +#[derive(Debug, Clone, Default)] +pub struct AgentSnapshotMemoryView { + pub history_len: usize, + pub current_text: String, + pub pending_tool_uses: Vec, +} + +impl From<&AgentMemoryState> for AgentSnapshotMemoryView { + fn from(state: &AgentMemoryState) -> Self { + Self { + history_len: state.transcript.len(), + current_text: state + .pending_turn + .as_ref() + .map(|pending| pending.current_text.clone()) + .unwrap_or_default(), + pending_tool_uses: state + .pending_turn + .as_ref() + .map(|pending| pending.pending_tool_uses.clone()) + .unwrap_or_default(), + } + } +} diff --git a/vendor/mentra/src/memory/journal/state.rs b/vendor/mentra/src/memory/journal/state.rs new file mode 100644 index 0000000..37abed9 --- /dev/null +++ b/vendor/mentra/src/memory/journal/state.rs @@ -0,0 +1,48 @@ +use serde::{Deserialize, Serialize}; + +use crate::{Message, agent::PendingToolUseSummary, transcript::AgentTranscript}; + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct AgentMemoryState { + #[serde(default, deserialize_with = "deserialize_transcript")] + pub transcript: AgentTranscript, + pub pending_turn: Option, + pub resumable_user_message: Option, + pub compaction: CompactionState, + pub revision: u64, + pub run: Option, +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct PendingTurnState { + pub current_text: String, + pub pending_tool_uses: Vec, +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct CompactionState; + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RunMemoryState { + pub run_id: String, + #[serde(default, deserialize_with = "deserialize_transcript")] + pub baseline_transcript: AgentTranscript, + pub assistant_committed: bool, +} + +fn deserialize_transcript<'de, D>(deserializer: D) -> Result +where + D: serde::Deserializer<'de>, +{ + #[derive(Deserialize)] + #[serde(untagged)] + enum TranscriptRepr { + Transcript(AgentTranscript), + Legacy(Vec), + } + + Ok(match TranscriptRepr::deserialize(deserializer)? { + TranscriptRepr::Transcript(transcript) => transcript, + TranscriptRepr::Legacy(messages) => AgentTranscript::from_messages(messages), + }) +} diff --git a/vendor/mentra/src/memory/journal/store.rs b/vendor/mentra/src/memory/journal/store.rs new file mode 100644 index 0000000..dad18a2 --- /dev/null +++ b/vendor/mentra/src/memory/journal/store.rs @@ -0,0 +1,16 @@ +use crate::{error::RuntimeError, runtime::RuntimeStore}; + +use super::state::AgentMemoryState; + +pub(crate) trait AgentMemoryStore: Send + Sync { + fn save_memory(&self, agent_id: &str, state: &AgentMemoryState) -> Result<(), RuntimeError>; +} + +impl AgentMemoryStore for T +where + T: RuntimeStore + ?Sized, +{ + fn save_memory(&self, agent_id: &str, state: &AgentMemoryState) -> Result<(), RuntimeError> { + self.save_agent_memory(agent_id, state) + } +} diff --git a/vendor/mentra/src/memory/journal/tests.rs b/vendor/mentra/src/memory/journal/tests.rs new file mode 100644 index 0000000..8169986 --- /dev/null +++ b/vendor/mentra/src/memory/journal/tests.rs @@ -0,0 +1,132 @@ +use std::{collections::BTreeMap, sync::Arc}; + +use serde_json::json; + +use crate::{ + AgentTranscript, ContentBlock, Message, TranscriptKind, + memory::journal::{AgentMemory, AgentMemoryState, CompactionOutcome, PendingTurnState}, + runtime::VolatileRuntimeStore, +}; + +#[test] +fn begin_run_commit_and_finish_persist_memory_state() { + let store = Arc::new(VolatileRuntimeStore::new()); + let mut memory = AgentMemory::new("agent-test", store, AgentMemoryState::default()); + + memory + .begin_run( + "run-1".to_string(), + Message::user(ContentBlock::text("hello")), + ) + .expect("begin run"); + assert_eq!(memory.transcript().len(), 1); + assert_eq!(memory.state().revision, 1); + assert_eq!( + memory.resumable_user_message(), + Some(&Message::user(ContentBlock::text("hello"))) + ); + + memory + .update_pending_turn(PendingTurnState::new("Hel".to_string(), Vec::new())) + .expect("update pending"); + assert_eq!(memory.snapshot_view().current_text, "Hel"); + + memory + .commit_assistant_message(Message::assistant(ContentBlock::text("done"))) + .expect("commit message"); + assert_eq!(memory.transcript().len(), 2); + assert!(memory.snapshot_view().current_text.is_empty()); + + memory.finish_run().expect("finish run"); + assert!(memory.state().run.is_none()); + assert!(memory.resumable_user_message().is_none()); +} + +#[test] +fn rollback_and_compaction_update_memory_state() { + let store = Arc::new(VolatileRuntimeStore::new()); + let mut memory = AgentMemory::new("agent-test", store, AgentMemoryState::default()); + + memory + .begin_run( + "run-1".to_string(), + Message::user(ContentBlock::text("hello")), + ) + .expect("begin run"); + memory + .update_pending_turn(PendingTurnState::new("partial".to_string(), Vec::new())) + .expect("pending"); + memory.rollback_failed_run().expect("rollback"); + assert!(memory.transcript().is_empty()); + assert_eq!( + memory.resumable_user_message(), + Some(&Message::user(ContentBlock::text("hello"))) + ); + + memory + .append_message(Message::user(ContentBlock::text("after"))) + .expect("append"); + let path = std::env::temp_dir().join("compacted.jsonl"); + memory + .compact(CompactionOutcome { + transcript_path: path.clone(), + transcript: AgentTranscript::from_messages(vec![Message::user(ContentBlock::text( + "summary", + ))]), + }) + .expect("compact"); + assert_eq!(memory.transcript().len(), 1); + let _ = path; +} + +// M3: `append_message_with_details` is an additive counterpart to +// `append_message` that behaves identically except for attaching metadata — +// proven directly against the in-process transcript here. The full +// persist/reload round-trip through the SQLite store (which additionally +// requires a real `agents` row, written by `Runtime::spawn`/`create_agent`, +// not just `AgentMemory` in isolation) is covered end-to-end in +// `agent::tests::runtime_resume::resumed_agent_keeps_tool_result_details_after_restart`. +#[test] +fn append_message_with_details_attaches_metadata_keyed_by_tool_use_id() { + let store = Arc::new(VolatileRuntimeStore::new()); + let mut memory = AgentMemory::new("agent-details", store, AgentMemoryState::default()); + + memory + .begin_run( + "run-1".to_string(), + Message::user(ContentBlock::text("run the details tool")), + ) + .expect("begin run"); + memory + .commit_assistant_message(Message::assistant(ContentBlock::ToolUse { + id: "call-1".to_string(), + name: "details_tool".to_string(), + input: json!({}), + })) + .expect("commit assistant tool call"); + + let details: BTreeMap = + BTreeMap::from([("call-1".to_string(), json!({ "secret": "shh", "n": 42 }))]); + memory + .append_message_with_details( + Message::user(ContentBlock::ToolResult { + tool_use_id: "call-1".to_string(), + content: "tool output".to_string().into(), + is_error: false, + }), + details.clone(), + ) + .expect("append with details"); + + let item = memory + .transcript() + .items() + .iter() + .find(|item| matches!(item.kind, TranscriptKind::ToolExchange { .. })) + .expect("transcript keeps the tool exchange item"); + assert_eq!(item.details(), Some(&details)); + assert_eq!( + item.detail("call-1"), + Some(&json!({ "secret": "shh", "n": 42 })) + ); +} diff --git a/vendor/mentra/src/provider.rs b/vendor/mentra/src/provider.rs new file mode 100644 index 0000000..2200f16 --- /dev/null +++ b/vendor/mentra/src/provider.rs @@ -0,0 +1,890 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use async_trait::async_trait; + +pub use mentra_provider::AnthropicRequestOptions; +pub use mentra_provider::AuthScheme; +pub use mentra_provider::BuiltinProvider; +pub use mentra_provider::CompactionInputItem; +pub use mentra_provider::CompactionRequest; +pub use mentra_provider::CompactionResponse; +pub use mentra_provider::ContentBlock; +pub use mentra_provider::ContentBlockDelta; +pub use mentra_provider::ContentBlockStart; +pub use mentra_provider::EmbeddingData; +pub use mentra_provider::EmbeddingModelInfo; +pub use mentra_provider::EmbeddingProvider; +pub use mentra_provider::EmbeddingRequest; +pub use mentra_provider::EmbeddingResponse; +pub use mentra_provider::EmbeddingUsage; +pub use mentra_provider::GeminiRequestOptions; +pub use mentra_provider::ImageSource; +pub use mentra_provider::MemorySummarizeOutput; +pub use mentra_provider::MemorySummarizeRequest; +pub use mentra_provider::MemorySummarizeResponse; +pub use mentra_provider::Message; +pub use mentra_provider::ModelInfo; +pub use mentra_provider::ModelSelector; +pub use mentra_provider::OpenAIRequestOptions; +pub use mentra_provider::ProviderCapabilities; +pub use mentra_provider::ProviderCredentials; +pub use mentra_provider::ProviderDefinition; +pub use mentra_provider::ProviderDescriptor; +pub use mentra_provider::ProviderError; +pub use mentra_provider::ProviderEvent; +pub use mentra_provider::ProviderEventStream; +pub use mentra_provider::ProviderId; +pub use mentra_provider::ProviderRequestOptions; +pub use mentra_provider::RawMemory; +pub use mentra_provider::RawMemoryMetadata; +pub use mentra_provider::ReasoningEffort; +pub use mentra_provider::ReasoningFormat; +pub use mentra_provider::ReasoningOptions; +pub use mentra_provider::ReasoningProvenance; +pub use mentra_provider::Request; +pub use mentra_provider::Response; +pub use mentra_provider::ResponsesRequestOptions; +pub use mentra_provider::ResponsesStateMode; +pub use mentra_provider::ResponsesTransport; +pub use mentra_provider::RetryPolicy; +pub use mentra_provider::Role; +pub use mentra_provider::TokenUsage; +pub use mentra_provider::ToolChoice; +pub use mentra_provider::ToolSearchMode; +pub use mentra_provider::WireApi; +pub use mentra_provider::collect_response_from_stream; +pub use mentra_provider::provider_event_stream_from_response; + +pub mod model { + pub use mentra_provider::AnthropicRequestOptions; + pub use mentra_provider::ContentBlock; + pub use mentra_provider::ContentBlockDelta; + pub use mentra_provider::ContentBlockStart; + pub use mentra_provider::ImageSource; + pub use mentra_provider::MemorySummarizeOutput; + pub use mentra_provider::MemorySummarizeRequest; + pub use mentra_provider::MemorySummarizeResponse; + pub use mentra_provider::Message; + pub use mentra_provider::ModelInfo; + pub use mentra_provider::OpenAIRequestOptions; + pub use mentra_provider::ProviderError; + pub use mentra_provider::ProviderEvent; + pub use mentra_provider::ProviderEventStream; + pub use mentra_provider::ProviderId; + pub use mentra_provider::ProviderRequestOptions; + pub use mentra_provider::RawMemory; + pub use mentra_provider::RawMemoryMetadata; + pub use mentra_provider::ReasoningEffort; + pub use mentra_provider::ReasoningFormat; + pub use mentra_provider::ReasoningOptions; + pub use mentra_provider::ReasoningProvenance; + pub use mentra_provider::Request; + pub use mentra_provider::Response; + pub use mentra_provider::ResponsesStateMode; + pub use mentra_provider::ResponsesTransport; + pub use mentra_provider::Role; + pub use mentra_provider::TokenUsage; + pub use mentra_provider::ToolChoice; + pub use mentra_provider::ToolSearchMode; + pub use mentra_provider::collect_response_from_stream; + pub use mentra_provider::provider_event_stream_from_response; +} + +/// Transport-neutral interface implemented by model providers. +#[async_trait] +pub trait Provider: Send + Sync { + /// Returns identifying metadata for the provider instance. + fn descriptor(&self) -> ProviderDescriptor; + + /// Returns feature flags supported by this provider instance. + fn capabilities(&self) -> ProviderCapabilities { + ProviderCapabilities::default() + } + + /// Lists models available from the provider. + async fn list_models(&self) -> Result, ProviderError>; + + /// Streams a model response for the given request. + async fn stream(&self, request: Request<'_>) -> Result; + + /// Sends a request and collects the full response in memory. + async fn send(&self, request: Request<'_>) -> Result { + collect_response_from_stream(self.stream(request).await?).await + } + + /// Compacts transcript history using a provider-native endpoint when supported. + async fn compact( + &self, + _request: CompactionRequest<'_>, + ) -> Result { + Err(ProviderError::UnsupportedCapability( + "history_compaction".to_string(), + )) + } + + /// Summarizes raw trace memories using a provider-native implementation when supported. + async fn summarize_memories( + &self, + _request: MemorySummarizeRequest<'_>, + ) -> Result { + Err(ProviderError::UnsupportedCapability( + "memory_summarization".to_string(), + )) + } +} + +#[derive(Default)] +pub struct ProviderRegistry { + default_provider: Option, + default_embedding_provider: Option, + providers: HashMap>, + embedding_providers: HashMap>, + /// The Responses transport this runtime's requests go out on, or `None` + /// when the runtime does not choose and each request's own options stand. + /// + /// It lives here rather than on the handle because a transport is a + /// property of the connection to a provider, and this is where a runtime + /// keeps those. It also means the choice travels the one path the builder + /// already hands to the handle at build time, instead of a field every + /// `with_*` reconstructor would have to remember to carry. + responses_transport: Option, +} + +impl ProviderRegistry { + pub(crate) fn register_builtin_provider( + &mut self, + id: BuiltinProvider, + api_key: impl Into, + ) -> Result<(), String> { + let api_key = api_key.into(); + let provider: Arc = match id { + BuiltinProvider::Anthropic => { + Arc::new(anthropic::AnthropicProvider::new(api_key.clone())) + } + BuiltinProvider::Gemini => Arc::new(gemini::GeminiProvider::new(api_key.clone())), + BuiltinProvider::OpenAI => Arc::new(openai::OpenAIProvider::new(api_key.clone())), + BuiltinProvider::OpenRouter => { + Arc::new(openrouter::OpenRouterProvider::new(api_key.clone())) + } + BuiltinProvider::Ollama => Arc::new(ollama::OllamaProvider::new()), + BuiltinProvider::LmStudio => Arc::new(lmstudio::LmStudioProvider::new()), + }; + + let provider_id: ProviderId = id.into(); + + if self.default_provider.is_none() { + self.default_provider = Some(provider_id.clone()); + } + + // Register embedding provider for providers that support it. + let ep: Option> = match id { + BuiltinProvider::OpenAI => Some(Arc::new(mentra_provider::responses::openai(api_key))), + BuiltinProvider::OpenRouter => { + Some(Arc::new(mentra_provider::responses::openrouter(api_key))) + } + BuiltinProvider::Ollama => Some(Arc::new(openai_compatible_embedding_provider( + id, + "http://127.0.0.1:11434/", + ))), + BuiltinProvider::LmStudio => Some(Arc::new(openai_compatible_embedding_provider( + id, + "http://127.0.0.1:1234/", + ))), + _ => None, + }; + if let Some(ep) = ep { + if self.default_embedding_provider.is_none() { + self.default_embedding_provider = Some(provider_id.clone()); + } + self.embedding_providers.insert(provider_id.clone(), ep); + } + + self.providers.insert(provider_id, provider); + Ok(()) + } + + pub(crate) fn register_provider_instance

(&mut self, provider: P) + where + P: Provider + 'static, + { + let descriptor = provider.descriptor(); + let id = descriptor.id; + + if self.default_provider.is_none() { + self.default_provider = Some(id.clone()); + } + + self.providers.insert(id, Arc::new(provider)); + } + + pub(crate) fn register_registered_provider

(&mut self, provider: P) + where + P: mentra_provider::Provider + 'static, + { + let descriptor = provider.descriptor(); + let id = descriptor.id; + + if self.default_provider.is_none() { + self.default_provider = Some(id.clone()); + } + + self.providers.insert(id, shared_provider(provider)); + } + + pub(crate) fn register_ollama(&mut self) { + self.register_provider_instance(ollama::OllamaProvider::new()); + } + + pub(crate) fn register_lmstudio(&mut self) { + self.register_provider_instance(lmstudio::LmStudioProvider::new()); + } + + pub(crate) fn get_provider(&self, id: Option<&ProviderId>) -> Option> { + match id { + Some(id) => self.providers.get(id).cloned(), + None => self + .default_provider + .as_ref() + .and_then(|id| self.providers.get(id).cloned()), + } + } + + /// Returns the default embedding provider, or `None` if no embedding-capable provider + /// has been registered. + /// + /// The default is the first embedding-capable provider registered. To look up a + /// specific provider use [`embedding_provider_for`]. + pub fn embedding_provider(&self) -> Option> { + self.default_embedding_provider + .as_ref() + .and_then(|id| self.embedding_providers.get(id)) + .map(Arc::clone) + .or_else(|| self.embedding_providers.values().next().map(Arc::clone)) + } + + /// Returns the embedding provider for a specific provider ID, or `None`. + pub fn embedding_provider_for(&self, id: &ProviderId) -> Option> { + self.embedding_providers.get(id).map(Arc::clone) + } + + pub(crate) fn descriptors(&self) -> Vec { + self.providers + .values() + .map(|provider| provider.descriptor()) + .collect() + } + + pub(crate) fn is_empty(&self) -> bool { + self.providers.is_empty() + } + + pub(crate) fn set_responses_transport(&mut self, transport: ResponsesTransport) { + self.responses_transport = Some(transport); + } + + pub(crate) fn responses_transport(&self) -> Option { + self.responses_transport + } +} + +/// Settles which Responses transport a request goes out on, and refuses one the +/// provider cannot serve. +/// +/// The runtime's choice, when it made one, replaces whatever the request's own +/// options carried: it is the connection-level answer, and a per-request one +/// that disagreed would mean two live opinions about a single socket. With no +/// runtime choice the request's own value stands, which is what every caller +/// had before a runtime could choose at all. +/// +/// A provider whose capabilities report no websocket support is refused rather +/// than quietly served over HTTP+SSE. The fallback is the tempting behavior and +/// the wrong one: asking for a transport is explicit, so answering on a +/// different one returns a stream nobody asked for and hides a misconfigured +/// runtime behind a working one — the same stance `stream_response` already +/// takes when the transport is not compiled in. +pub(crate) fn select_responses_transport( + provider: &dyn Provider, + chosen: Option, + options: &mut ProviderRequestOptions, +) -> Result<(), crate::error::RuntimeError> { + if let Some(transport) = chosen { + options.responses.transport = transport; + } + + if options.responses.transport != ResponsesTransport::WebSocket + || provider.capabilities().supports_websockets + { + return Ok(()); + } + + let descriptor = provider.descriptor(); + let name = descriptor + .display_name + .unwrap_or_else(|| descriptor.id.as_str().to_string()); + Err(crate::error::RuntimeError::OperationDenied(format!( + "provider '{name}' does not serve the Responses websocket transport; \ + select ResponsesTransport::HttpSse or register a provider that does \ + — answering over HTTP+SSE would return a transport nobody asked for" + ))) +} + +fn shared_provider

(provider: P) -> Arc +where + P: mentra_provider::Provider + 'static, +{ + Arc::new(SharedProviderProxy { inner: provider }) +} + +/// Builds a `ResponsesProvider` (with no credentials) for OpenAI-compatible +/// local providers (Ollama, LmStudio) so they can be used as embedding providers. +fn openai_compatible_embedding_provider( + builtin: BuiltinProvider, + base_url: &str, +) -> mentra_provider::responses::ResponsesProvider { + use mentra_provider::AuthScheme; + use mentra_provider::ProviderCapabilities; + use mentra_provider::RetryPolicy; + use mentra_provider::WireApi; + use std::collections::HashMap; + + let mut definition = ProviderDefinition::new(builtin); + definition.wire_api = WireApi::Responses; + definition.auth_scheme = AuthScheme::None; + definition.capabilities = ProviderCapabilities { + supports_model_listing: true, + supports_streaming: true, + supports_websockets: false, + supports_tool_calls: true, + supports_images: true, + supports_history_compaction: false, + supports_memory_summarization: false, + supports_deferred_tools: false, + supports_hosted_tool_search: false, + supports_hosted_web_search: false, + supports_image_generation: false, + supports_reasoning_effort: false, + reports_reasoning_tokens: false, + reports_thoughts_tokens: false, + supports_structured_tool_results: false, + supports_embeddings: true, + }; + definition.base_url = Some(base_url.to_string()); + definition.headers = Some(HashMap::new()); + definition.retry = RetryPolicy::default(); + mentra_provider::responses::ResponsesProvider::new(definition, NoCredentialsSource) +} + +#[derive(Clone)] +struct NoCredentialsSource; + +#[async_trait] +impl mentra_provider::CredentialSource for NoCredentialsSource { + async fn credentials( + &self, + ) -> Result { + Ok(mentra_provider::ProviderCredentials::default()) + } +} + +struct SharedProviderProxy

{ + inner: P, +} + +#[async_trait] +impl

Provider for SharedProviderProxy

+where + P: mentra_provider::Provider + 'static, +{ + fn descriptor(&self) -> ProviderDescriptor { + self.inner.descriptor() + } + + fn capabilities(&self) -> ProviderCapabilities { + self.inner.definition().capabilities + } + + async fn list_models(&self) -> Result, ProviderError> { + self.inner.list_models().await + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.inner.stream(request).await + } + + async fn compact( + &self, + request: CompactionRequest<'_>, + ) -> Result { + self.inner.compact(request).await + } + + async fn summarize_memories( + &self, + request: MemorySummarizeRequest<'_>, + ) -> Result { + self.inner.summarize_memories(request).await + } +} + +pub mod openai { + use std::collections::HashMap; + use std::sync::Arc; + + use async_trait::async_trait; + + use super::AuthScheme; + use super::BuiltinProvider; + use super::CompactionRequest; + use super::CompactionResponse; + use super::Provider; + use super::ProviderCapabilities; + use super::ProviderDefinition; + use super::ProviderDescriptor; + use super::ProviderError; + use super::ProviderEventStream; + use super::Request; + use super::RetryPolicy; + use super::WireApi; + use super::shared_provider; + + use crate::provider::model::ModelInfo; + + /// Supplies OpenAI API credentials on demand. + #[async_trait] + pub trait OpenAICredentialSource: Send + Sync { + async fn api_key(&self) -> Result; + } + + #[derive(Clone)] + pub struct OpenAIProvider { + inner: Arc, + } + + impl OpenAIProvider { + pub fn new(api_key: impl Into) -> Self { + Self { + inner: shared_provider(mentra_provider::responses::openai(api_key)), + } + } + + pub(crate) fn openai_compatible( + provider: BuiltinProvider, + display_name: &'static str, + description: &'static str, + base_url: &str, + ) -> Self { + let mut definition = ProviderDefinition::new(provider); + definition.descriptor.display_name = Some(display_name.to_string()); + definition.descriptor.description = Some(description.to_string()); + definition.wire_api = WireApi::Responses; + definition.auth_scheme = AuthScheme::None; + definition.capabilities = ProviderCapabilities { + supports_model_listing: true, + supports_streaming: true, + supports_websockets: false, + supports_tool_calls: true, + supports_images: true, + supports_history_compaction: false, + supports_memory_summarization: false, + supports_deferred_tools: false, + supports_hosted_tool_search: false, + supports_hosted_web_search: false, + supports_image_generation: false, + supports_reasoning_effort: false, + reports_reasoning_tokens: false, + reports_thoughts_tokens: false, + supports_structured_tool_results: false, + supports_embeddings: true, + }; + definition.base_url = Some(base_url.to_string()); + definition.headers = Some(HashMap::new()); + definition.retry = RetryPolicy::default(); + + let provider = mentra_provider::responses::ResponsesProvider::new( + definition, + super::NoCredentialsSource, + ); + Self { + inner: shared_provider(provider), + } + } + + pub fn with_credential_source(source: impl OpenAICredentialSource + 'static) -> Self { + Self::with_shared_credential_source(Arc::new(source)) + } + + pub fn with_shared_credential_source(source: Arc) -> Self { + let provider = mentra_provider::responses::openai_with_credential_source( + OpenAICredentialAdapter { source }, + ); + Self { + inner: shared_provider(provider), + } + } + } + + #[async_trait] + impl Provider for OpenAIProvider { + fn descriptor(&self) -> ProviderDescriptor { + self.inner.descriptor() + } + + fn capabilities(&self) -> ProviderCapabilities { + self.inner.capabilities() + } + + async fn list_models(&self) -> Result, ProviderError> { + self.inner.list_models().await + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.inner.stream(request).await + } + async fn compact( + &self, + request: CompactionRequest<'_>, + ) -> Result { + self.inner.compact(request).await + } + + async fn summarize_memories( + &self, + request: super::MemorySummarizeRequest<'_>, + ) -> Result { + self.inner.summarize_memories(request).await + } + } + + #[derive(Clone)] + struct OpenAICredentialAdapter { + source: Arc, + } + + #[async_trait] + impl mentra_provider::CredentialSource for OpenAICredentialAdapter { + async fn credentials( + &self, + ) -> Result { + let api_key = self + .source + .api_key() + .await + .map_err(mentra_provider::ProviderError::InvalidRequest)?; + + Ok(mentra_provider::ProviderCredentials { + bearer_token: Some(api_key), + account_id: None, + headers: Default::default(), + }) + } + } +} + +pub mod openrouter { + use std::sync::Arc; + + use async_trait::async_trait; + + use super::CompactionRequest; + use super::CompactionResponse; + use super::Provider; + use super::ProviderCapabilities; + use super::ProviderDescriptor; + use super::ProviderError; + use super::ProviderEventStream; + use super::Request; + use super::shared_provider; + use crate::provider::model::ModelInfo; + + #[derive(Clone)] + pub struct OpenRouterProvider { + inner: Arc, + } + + impl OpenRouterProvider { + pub fn new(api_key: impl Into) -> Self { + Self { + inner: shared_provider(mentra_provider::responses::openrouter(api_key)), + } + } + } + + #[async_trait] + impl Provider for OpenRouterProvider { + fn descriptor(&self) -> ProviderDescriptor { + self.inner.descriptor() + } + + fn capabilities(&self) -> ProviderCapabilities { + self.inner.capabilities() + } + + async fn list_models(&self) -> Result, ProviderError> { + self.inner.list_models().await + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.inner.stream(request).await + } + async fn compact( + &self, + request: CompactionRequest<'_>, + ) -> Result { + self.inner.compact(request).await + } + + async fn summarize_memories( + &self, + request: super::MemorySummarizeRequest<'_>, + ) -> Result { + self.inner.summarize_memories(request).await + } + } +} + +pub mod anthropic { + use std::sync::Arc; + + use async_trait::async_trait; + + use super::Provider; + use super::ProviderDescriptor; + use super::ProviderError; + use super::ProviderEventStream; + use super::Request; + use super::shared_provider; + use crate::provider::model::ModelInfo; + + #[derive(Clone)] + pub struct AnthropicProvider { + inner: Arc, + } + + impl AnthropicProvider { + pub fn new(api_key: impl Into) -> Self { + Self { + inner: shared_provider(mentra_provider::anthropic::AnthropicProvider::new(api_key)), + } + } + } + + #[async_trait] + impl Provider for AnthropicProvider { + fn descriptor(&self) -> ProviderDescriptor { + self.inner.descriptor() + } + + fn capabilities(&self) -> super::ProviderCapabilities { + self.inner.capabilities() + } + + async fn list_models(&self) -> Result, ProviderError> { + self.inner.list_models().await + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.inner.stream(request).await + } + + async fn summarize_memories( + &self, + request: super::MemorySummarizeRequest<'_>, + ) -> Result { + self.inner.summarize_memories(request).await + } + } +} + +pub mod gemini { + use std::sync::Arc; + + use async_trait::async_trait; + + use super::Provider; + use super::ProviderDescriptor; + use super::ProviderError; + use super::ProviderEventStream; + use super::Request; + use super::shared_provider; + use crate::provider::model::ModelInfo; + + #[derive(Clone)] + pub struct GeminiProvider { + inner: Arc, + } + + impl GeminiProvider { + pub fn new(api_key: impl Into) -> Self { + Self { + inner: shared_provider(mentra_provider::gemini::GeminiProvider::new(api_key)), + } + } + } + + #[async_trait] + impl Provider for GeminiProvider { + fn descriptor(&self) -> ProviderDescriptor { + self.inner.descriptor() + } + + fn capabilities(&self) -> super::ProviderCapabilities { + self.inner.capabilities() + } + + async fn list_models(&self) -> Result, ProviderError> { + self.inner.list_models().await + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.inner.stream(request).await + } + + async fn summarize_memories( + &self, + request: super::MemorySummarizeRequest<'_>, + ) -> Result { + self.inner.summarize_memories(request).await + } + } +} + +pub mod ollama { + use std::sync::Arc; + + use async_trait::async_trait; + + use super::BuiltinProvider; + use super::Provider; + use super::ProviderDescriptor; + use super::ProviderError; + use super::ProviderEventStream; + use super::Request; + use crate::provider::model::ModelInfo; + + const DEFAULT_BASE_URL: &str = "http://127.0.0.1:11434/"; + + #[derive(Clone)] + pub struct OllamaProvider { + inner: Arc, + } + + impl OllamaProvider { + pub fn new() -> Self { + Self::with_base_url(DEFAULT_BASE_URL) + } + + pub fn with_base_url(base_url: impl AsRef) -> Self { + Self { + inner: Arc::new(super::openai::OpenAIProvider::openai_compatible( + BuiltinProvider::Ollama, + "Ollama", + "Ollama OpenAI-compatible Responses API provider", + base_url.as_ref(), + )), + } + } + } + + impl Default for OllamaProvider { + fn default() -> Self { + Self::new() + } + } + + #[async_trait] + impl Provider for OllamaProvider { + fn descriptor(&self) -> ProviderDescriptor { + self.inner.descriptor() + } + + fn capabilities(&self) -> super::ProviderCapabilities { + self.inner.capabilities() + } + + async fn list_models(&self) -> Result, ProviderError> { + self.inner.list_models().await + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.inner.stream(request).await + } + + async fn summarize_memories( + &self, + request: super::MemorySummarizeRequest<'_>, + ) -> Result { + self.inner.summarize_memories(request).await + } + } +} + +pub mod lmstudio { + use std::sync::Arc; + + use async_trait::async_trait; + + use super::BuiltinProvider; + use super::Provider; + use super::ProviderDescriptor; + use super::ProviderError; + use super::ProviderEventStream; + use super::Request; + use crate::provider::model::ModelInfo; + + const DEFAULT_BASE_URL: &str = "http://127.0.0.1:1234/"; + + #[derive(Clone)] + pub struct LmStudioProvider { + inner: Arc, + } + + impl LmStudioProvider { + pub fn new() -> Self { + Self::with_base_url(DEFAULT_BASE_URL) + } + + pub fn with_base_url(base_url: impl AsRef) -> Self { + Self { + inner: Arc::new(super::openai::OpenAIProvider::openai_compatible( + BuiltinProvider::LmStudio, + "LM Studio", + "LM Studio OpenAI-compatible Responses API provider", + base_url.as_ref(), + )), + } + } + } + + impl Default for LmStudioProvider { + fn default() -> Self { + Self::new() + } + } + + #[async_trait] + impl Provider for LmStudioProvider { + fn descriptor(&self) -> ProviderDescriptor { + self.inner.descriptor() + } + + fn capabilities(&self) -> super::ProviderCapabilities { + self.inner.capabilities() + } + + async fn list_models(&self) -> Result, ProviderError> { + self.inner.list_models().await + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.inner.stream(request).await + } + + async fn summarize_memories( + &self, + request: super::MemorySummarizeRequest<'_>, + ) -> Result { + self.inner.summarize_memories(request).await + } + } +} diff --git a/vendor/mentra/src/runtime.rs b/vendor/mentra/src/runtime.rs new file mode 100644 index 0000000..5523bba --- /dev/null +++ b/vendor/mentra/src/runtime.rs @@ -0,0 +1,638 @@ +mod builder; +pub(crate) mod control; +mod error; +pub(crate) mod handle; +mod hybrid_store; +mod intrinsic; +mod skill; +mod store; +pub(crate) mod task; +mod task_board; +mod volatile_store; + +use std::{any::Any, path::Path, sync::Arc}; + +use tokio::sync::broadcast; + +use crate::{ + agent::{Agent, AgentConfig, AgentSpawnOptions, AgentStatus}, + provider::{Provider, ProviderRegistry}, + session::{ + Session, SessionEvent, SessionId, SessionMetadata, + permission::{PendingPermissionStore, SessionToolAuthorizer}, + }, + tool::ExecutableTool, +}; +use mentra_provider::{BuiltinProvider, ModelInfo, ModelSelector, ProviderDescriptor, ProviderId}; + +pub use builder::RuntimeBuilder; +pub use control::sandbox::{ExecutionEnvironment, detect_environment}; +pub use control::{ + AuditHook, AuditLogHook, CancellationFlag, CancellationToken, CommandOutput, CommandRequest, + CommandSpec, EarlyEnd, ExecOutput, HookDecision, LocalRuntimeExecutor, PreExecutionContext, + PreExecutionHook, PreExecutionHooks, ProviderRetry, RunOptions, RuntimeExecutor, RuntimeHook, + RuntimeHookEvent, RuntimeHooks, RuntimePolicy, ShellValidationMode, + is_transient_provider_error, is_transient_runtime_error, +}; +pub use error::{ErrorCategory, RuntimeError}; +pub(crate) use handle::RuntimeHandle; +pub use hybrid_store::HybridRuntimeStore; +pub(crate) use intrinsic::RuntimeIntrinsicTool; +pub use skill::{SkillInfo, SkillLoadError}; +pub use store::{ + AgentStore, AuditStore, LeaseStore, PermissionRuleStore, RunStore, RuntimeStore, + SqliteRuntimeStore, TaskStore, +}; +pub(crate) use store::{LoadedAgentState, PersistedAgentRecord, TaskStateSnapshot}; +pub(crate) use task::TaskIntrinsicTool; +pub use task::{TaskItem, TaskStatus}; +pub use task_board::{NewTask, TaskBoard, TaskBoardError, TaskPatch}; +pub use volatile_store::VolatileRuntimeStore; + +/// Entry point for configuring providers, tools, and agent lifecycles. +/// +/// A runtime composes four main subsystems: +/// - execution: providers, policies, hooks, and command execution +/// - persistence: agent state, runs, tasks, leases, and memory +/// - tooling: registered tools, skills, and app context +/// - collaboration: persistent teams and background task coordination +pub struct Runtime { + handle: RuntimeHandle, + provider_registry: Arc>, + pub(crate) mcp_servers: Vec, +} + +/// How one configured MCP server fared during +/// [`build_async`](RuntimeBuilder::build_async). +/// +/// A server that fails to connect leaves the runtime in degraded mode rather +/// than failing the build — one unreachable server should not sink a session. +/// This is how a host finds out which ones are actually live, so it can say so +/// instead of leaving a user to wonder why a tool is missing. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct McpServerSummary { + pub name: String, + /// Tools this server contributed. Zero when it failed. + pub tools: usize, + /// Why it did not connect, when it did not. + pub error: Option, +} + +impl McpServerSummary { + pub fn connected(&self) -> bool { + self.error.is_none() + } +} + +/// Read-only summary of a persisted agent record for a runtime identifier. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PersistedAgentSummary { + pub id: String, + pub runtime_identifier: String, + pub name: String, + pub is_teammate: bool, + pub status: AgentStatus, + pub history_len: usize, +} + +impl Runtime { + /// Returns a builder with Mentra's builtin tools enabled. + pub fn builder() -> RuntimeBuilder { + RuntimeBuilder::new(true) + } + + /// Returns a builder with no builtin tools registered. + pub fn empty_builder() -> RuntimeBuilder { + RuntimeBuilder::new(false) + } + + /// Registers a custom tool on the runtime after construction. + pub fn register_tool(&self, tool: T) + where + T: ExecutableTool + 'static, + { + self.handle.register_tool(tool); + } + + /// Returns descriptors for registered tools in a deterministic order. + pub fn tools(&self) -> Vec { + let tool_names = self + .handle + .tools() + .iter() + .map(|tool| tool.name.clone()) + .collect::>(); + let mut tools = tool_names + .into_iter() + .filter_map(|name| self.handle.get_tool_descriptor(&name)) + .collect::>(); + tools.sort_by(|left, right| left.provider.name.cmp(&right.provider.name)); + tools + } + + /// Returns the descriptor for a registered tool by name. + pub fn tool_descriptor(&self, name: &str) -> Option { + self.handle.get_tool_descriptor(name) + } + + /// Registers typed application state that tools can retrieve from their context. + pub fn register_context(&self, context: Arc) { + self.handle.register_app_context(context); + } + + /// Returns typed application state previously registered on this runtime. + pub fn app_context(&self) -> Result, String> + where + T: Any + Send + Sync + 'static, + { + self.handle.app_context::() + } + + /// Registers a skills directory and enables the builtin `load_skill` tool. + /// + /// Additive: calling this again adds a second root rather than replacing + /// the first, and a name already registered wins. Register the most + /// specific root first. + pub fn register_skills_dir(&self, path: impl AsRef) -> Result<(), SkillLoadError> { + self.handle + .register_skill_loader(skill::SkillLoader::from_dir(path)?); + Ok(()) + } + + /// Registers several skills directories at once, strongest first. + /// + /// Equivalent to calling [`register_skills_dir`](Self::register_skills_dir) + /// for each in order: a skill defined in an earlier root shadows the same + /// name in a later one, so a project root can override a personal one. + /// Within a single root a repeated name is still an error. + /// + /// Registration stops at the first unreadable root, leaving the roots + /// before it registered. + pub fn register_skills_dirs(&self, paths: I) -> Result<(), SkillLoadError> + where + I: IntoIterator, + P: AsRef, + { + for path in paths { + self.register_skills_dir(path)?; + } + Ok(()) + } + + /// Every loaded skill, name-ordered, with its description and source path + /// but not its body. + pub fn skills(&self) -> Vec { + self.handle.skills() + } + + /// How each configured MCP server fared while the runtime was built. + /// + /// Empty when none were configured, or when the runtime came from + /// [`build`](RuntimeBuilder::build), which refuses to be given any. + /// A failed server is present with its error rather than absent: a host + /// telling a user which tools they have needs to name what is missing. + pub fn mcp_servers(&self) -> &[McpServerSummary] { + &self.mcp_servers + } + + /// Returns a lead-privileged task-board view for `namespace`. + /// + /// The namespace is an opaque store key; no directory is created. Reads are + /// live and every mutation passes through the same validation and + /// transactional store path as the builtin task tools. + pub fn task_board(&self, namespace: impl AsRef) -> TaskBoard { + TaskBoard::lead(self.handle.clone(), namespace.as_ref().to_path_buf()) + } + + /// Spawns a new agent with the default [`AgentConfig`]. + pub fn spawn(&self, name: impl Into, model: ModelInfo) -> Result { + self.spawn_with_config(name, model, AgentConfig::default()) + } + + /// Spawns a new agent with an explicit configuration. + pub fn spawn_with_config( + &self, + name: impl Into, + model: ModelInfo, + config: AgentConfig, + ) -> Result { + Agent::new( + self.handle.clone(), + model.id, + name.into(), + config, + self.provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(Some(&model.provider)) + .ok_or_else(|| RuntimeError::ProviderNotFound(Some(model.provider.clone())))?, + AgentSpawnOptions::default(), + ) + } + + /// Restores a previously persisted agent by identifier. + pub fn resume_agent(&self, agent_id: &str) -> Result { + let Some(state) = self.handle.store().load_agent(agent_id)? else { + return Err(RuntimeError::Store(format!( + "No persisted agent with id '{agent_id}'" + ))); + }; + let provider = self + .provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(Some(&state.record.provider_id)) + .ok_or_else(|| { + RuntimeError::ProviderNotFound(Some(state.record.provider_id.clone())) + })?; + Agent::from_loaded(self.handle.clone(), state, provider) + } + + /// Restores every persisted agent that belongs to the provided runtime identifier. + pub fn resume(&self, runtime_identifier: &str) -> Result, RuntimeError> { + let states = self + .handle + .store() + .list_agents_by_runtime(runtime_identifier)?; + let mut agents = Vec::new(); + for state in states { + let provider = self + .provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(Some(&state.record.provider_id)) + .ok_or_else(|| { + RuntimeError::ProviderNotFound(Some(state.record.provider_id.clone())) + })?; + let agent = Agent::from_loaded(self.handle.clone(), state, provider)?; + if agent.is_teammate() { + agent.revive_teammate_actor()?; + } else { + agents.push(agent); + } + } + Ok(agents) + } + + /// Lists persisted agents for a runtime identifier without reviving them. + pub fn list_persisted_agents( + &self, + runtime_identifier: &str, + ) -> Result, RuntimeError> { + self.handle + .store() + .list_agents_by_runtime(runtime_identifier) + .map(|states| { + states + .into_iter() + .map(|state| PersistedAgentSummary { + id: state.record.id, + runtime_identifier: state.record.runtime_identifier, + name: state.record.name, + is_teammate: state.record.teammate_identity.is_some(), + status: state.record.status, + history_len: state.memory.transcript.len(), + }) + .collect() + }) + } + + /// Restores every persisted agent known to the runtime store. + pub fn resume_all(&self) -> Result, RuntimeError> { + let states = self.handle.store().list_agents()?; + let mut agents = Vec::new(); + for state in states { + let provider = self + .provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(Some(&state.record.provider_id)) + .ok_or_else(|| { + RuntimeError::ProviderNotFound(Some(state.record.provider_id.clone())) + })?; + agents.push(Agent::from_loaded(self.handle.clone(), state, provider)?); + } + Ok(agents) + } +} + +impl Runtime { + /// Returns descriptors for registered providers. + pub fn providers(&self) -> Vec { + self.provider_registry + .read() + .expect("provider registry poisoned") + .descriptors() + } + + /// The Responses transport this runtime chose for every request it makes, + /// or `None` when it left the choice to each request's own options — which + /// is HTTP+SSE unless a host set otherwise. + /// + /// The reader for + /// [`RuntimeBuilder::with_responses_transport`](crate::runtime::RuntimeBuilder::with_responses_transport). + /// A transport is otherwise the one piece of a runtime's configuration + /// nothing can observe: a registered tool shows up in + /// [`tools`](Self::tools), a provider in [`providers`](Self::providers), + /// but a transport reaches only the requests the runtime sends. That makes + /// the wiring between a host's choice and this runtime untestable except by + /// running a turn against a provider that records what it was handed — and + /// leaves a host that wants to report its own configuration with no way to + /// ask. + pub fn responses_transport(&self) -> Option { + self.provider_registry + .read() + .expect("provider registry poisoned") + .responses_transport() + } + + /// Registers a builtin provider from an API key. + pub fn register_provider( + &mut self, + id: BuiltinProvider, + api_key: impl Into, + ) -> Result<(), String> { + self.provider_registry + .write() + .expect("provider registry poisoned") + .register_builtin_provider(id, api_key) + } + + /// Registers the local Ollama provider using its default OpenAI-compatible endpoint. + pub fn register_ollama(&mut self) { + self.provider_registry + .write() + .expect("provider registry poisoned") + .register_ollama(); + } + + /// Registers the local LM Studio provider using its default OpenAI-compatible endpoint. + pub fn register_lmstudio(&mut self) { + self.provider_registry + .write() + .expect("provider registry poisoned") + .register_lmstudio(); + } + + /// Registers a custom runtime provider implementation. + /// + /// This is the supported seam for injecting a scripted provider in tests or + /// embedding Mentra on top of a custom transport. + /// + /// ```rust,no_run + /// use async_trait::async_trait; + /// use mentra::{BuiltinProvider, ModelInfo, ProviderDescriptor, Runtime}; + /// use mentra::error::{ProviderError, RuntimeError}; + /// use mentra::provider::{Provider, ProviderEventStream, Request}; + /// use tokio::sync::mpsc; + /// + /// struct TestProvider; + /// + /// #[async_trait] + /// impl Provider for TestProvider { + /// fn descriptor(&self) -> ProviderDescriptor { + /// ProviderDescriptor::new(BuiltinProvider::Anthropic) + /// } + /// + /// async fn list_models(&self) -> Result, ProviderError> { + /// Ok(vec![ModelInfo::new("test-model", BuiltinProvider::Anthropic)]) + /// } + /// + /// async fn stream( + /// &self, + /// _request: Request<'_>, + /// ) -> Result { + /// let (_tx, rx) = mpsc::unbounded_channel(); + /// Ok(rx) + /// } + /// } + /// + /// let mut runtime = Runtime::empty_builder() + /// .with_provider(BuiltinProvider::Anthropic, "placeholder") + /// .build()?; + /// runtime.register_provider_instance(TestProvider); + /// # Ok::<(), RuntimeError>(()) + /// ``` + pub fn register_provider_instance

(&mut self, provider: P) + where + P: Provider + 'static, + { + self.provider_registry + .write() + .expect("provider registry poisoned") + .register_provider_instance(provider); + } + + /// Registers a provider-core instance built from `mentra::provider_core`. + /// + /// Use this when you want Mentra's runtime with a customized provider + /// definition, such as a custom OpenAI-compatible or Anthropic-compatible + /// base URL. + pub fn register_registered_provider

(&mut self, provider: P) + where + P: mentra_provider::Provider + 'static, + { + self.provider_registry + .write() + .expect("provider registry poisoned") + .register_registered_provider(provider); + } + + /// Lists models for a specific provider, or the default provider when omitted. + pub async fn list_models( + &self, + provider: Option<&ProviderId>, + ) -> Result, RuntimeError> { + let provider = self + .provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(provider) + .ok_or_else(|| RuntimeError::ProviderNotFound(provider.cloned()))?; + + provider + .list_models() + .await + .map_err(RuntimeError::FailedToListModels) + } + + /// Resolves a model for a registered provider using a deterministic selection strategy. + pub async fn resolve_model( + &self, + provider: impl Into, + selector: ModelSelector, + ) -> Result { + let provider = provider.into(); + if self + .provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(Some(&provider)) + .is_none() + { + return Err(RuntimeError::ProviderNotFound(Some(provider))); + } + + match selector { + ModelSelector::Id(id) => Ok(ModelInfo::new(id, provider)), + ModelSelector::NewestAvailable => { + let mut models = self.list_models(Some(&provider)).await?; + models.sort_by(|left, right| { + right + .created_at + .cmp(&left.created_at) + .then_with(|| left.id.cmp(&right.id)) + }); + models + .into_iter() + .next() + .ok_or(RuntimeError::NoModelsAvailable(provider)) + } + } + } +} + +// -- Session lifecycle methods -- + +impl Runtime { + /// Creates a new session wrapping a freshly spawned agent with default config. + pub fn create_session( + &self, + name: impl Into, + model: ModelInfo, + ) -> Result { + self.create_session_with_config(name, model, AgentConfig::default()) + } + + /// Creates a new session wrapping a freshly spawned agent with explicit config. + /// + /// Convenience wrapper around [`create_session_full`](Self::create_session_full) that + /// passes `None` for `project_id`. + pub fn create_session_with_config( + &self, + name: impl Into, + model: ModelInfo, + config: AgentConfig, + ) -> Result { + self.create_session_full(name, model, config, None) + } + + /// Creates a new session wrapping a freshly spawned agent with explicit config and + /// an optional project identifier. + /// + /// The `project_id` is threaded into the [`SessionPermissionHandle`] so that + /// permission rules are scoped to the project when a [`PermissionRuleStore`] is + /// attached. + pub fn create_session_full( + &self, + name: impl Into, + model: ModelInfo, + config: AgentConfig, + project_id: Option, + ) -> Result { + let name = name.into(); + let session_id = SessionId::new(); + let metadata = SessionMetadata::new(session_id.clone(), &name, &model.id); + let (event_tx, _) = broadcast::channel(512); + let rule_store = crate::session::RuleStore::new(); + let pending_permissions = PendingPermissionStore::new(); + let session_handle = + self.handle + .with_tool_authorizer(Arc::new(SessionToolAuthorizer::new( + self.handle.execution.tool_authorizer.clone(), + event_tx.clone(), + pending_permissions.clone(), + rule_store.clone(), + ))); + let provider = self + .provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(Some(&model.provider)) + .ok_or_else(|| RuntimeError::ProviderNotFound(Some(model.provider.clone())))?; + let agent = Agent::new( + session_handle, + model.id.clone(), + name.clone(), + config, + provider, + AgentSpawnOptions::default(), + )?; + let mut session = Session::new_with_parts( + session_id.clone(), + metadata, + agent, + event_tx, + rule_store, + pending_permissions, + project_id, + ); + + // Emit the initial SessionStarted event. + let started = SessionEvent::SessionStarted { session_id }; + // Subscribe briefly just to ensure the event is broadcast. + let _rx = session.subscribe(); + // Use the internal emit path via a helper on Session. + session.emit_started(started); + + Ok(session) + } + + /// Resumes a previously persisted agent and wraps it in a session. + /// + /// Convenience wrapper around [`resume_session_with_project`](Self::resume_session_with_project) + /// that passes `None` for `project_id`. + pub fn resume_session(&self, agent_id: &str) -> Result { + self.resume_session_with_project(agent_id, None) + } + + /// Resumes a previously persisted agent, wraps it in a session, and associates + /// the session with an optional project identifier. + /// + /// The `project_id` is threaded into the [`SessionPermissionHandle`] so that + /// permission rules are scoped to the project when a [`PermissionRuleStore`] is + /// attached. + pub fn resume_session_with_project( + &self, + agent_id: &str, + project_id: Option, + ) -> Result { + let session_id = SessionId::new(); + let (event_tx, _) = broadcast::channel(512); + let rule_store = crate::session::RuleStore::new(); + let pending_permissions = PendingPermissionStore::new(); + let session_handle = + self.handle + .with_tool_authorizer(Arc::new(SessionToolAuthorizer::new( + self.handle.execution.tool_authorizer.clone(), + event_tx.clone(), + pending_permissions.clone(), + rule_store.clone(), + ))); + let Some(state) = self.handle.store().load_agent(agent_id)? else { + return Err(RuntimeError::Store(format!( + "No persisted agent with id '{agent_id}'" + ))); + }; + let provider = self + .provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(Some(&state.record.provider_id)) + .ok_or_else(|| { + RuntimeError::ProviderNotFound(Some(state.record.provider_id.clone())) + })?; + let agent = Agent::from_loaded(session_handle, state, provider)?; + let metadata = SessionMetadata::new(session_id.clone(), agent.name(), agent.model()); + let session = Session::new_with_parts( + session_id, + metadata, + agent, + event_tx, + rule_store, + pending_permissions, + project_id, + ); + Ok(session) + } +} diff --git a/vendor/mentra/src/runtime/builder.rs b/vendor/mentra/src/runtime/builder.rs new file mode 100644 index 0000000..e0d2f62 --- /dev/null +++ b/vendor/mentra/src/runtime/builder.rs @@ -0,0 +1,727 @@ +use std::{any::Any, path::Path, sync::Arc}; + +use crate::{ + compaction::CompactionEngine, + mcp::{McpManager, McpServerConfig, McpSseServerConfig}, + provider::{Provider, ProviderRegistry, ResponsesTransport}, + runtime::{ + RuntimeExecutor, RuntimeHandle, RuntimeHook, RuntimeHooks, RuntimePolicy, RuntimeStore, + control::PreExecutionHook, error::RuntimeError, skill::SkillLoadError, + }, + tool::{ExecutableTool, FileToolProfile, ToolAuthorizer}, +}; +use mentra_provider::BuiltinProvider; + +use super::skill::SkillLoader; +use super::{McpServerSummary, Runtime}; + +/// An MCP server to connect to during build, and how to reach it. +/// +/// This is internal so that the two public registration methods keep taking +/// their own configuration types: [`McpServerConfig`] stays the stdio +/// configuration and callers never gain a transport field to fill in. +enum McpRegistration { + Stdio(Box), + Sse(Box), +} + +impl McpRegistration { + /// The configured server name, used for diagnostics. + fn name(&self) -> &str { + match self { + Self::Stdio(config) => &config.name, + Self::Sse(config) => &config.name, + } + } +} + +/// Builder for constructing a [`Runtime`] with providers, tools, and policies. +pub struct RuntimeBuilder { + handle: RuntimeHandle, + provider_registry: ProviderRegistry, + mcp_configs: Vec, +} + +impl RuntimeBuilder { + /// Creates a builder with Mentra's builtin tools enabled. + pub fn new(runtime_intrinsics_enabled: bool) -> Self { + Self { + handle: RuntimeHandle::new(runtime_intrinsics_enabled), + provider_registry: ProviderRegistry::default(), + mcp_configs: Vec::new(), + } + } + + /// Registers a custom tool. + pub fn with_tool(self, tool: T) -> Self + where + T: ExecutableTool + 'static, + { + self.handle.register_tool(tool); + self + } + + /// Reconfigures the eagerly registered builtin file-tool surface. + /// + /// The default is [`FileToolProfile::Batched`], preserving the historical + /// `files` tool. This method also works with [`Runtime::empty_builder`] to + /// opt into only the selected file tools. + pub fn with_file_tools(self, profile: FileToolProfile) -> Self { + self.handle.configure_file_tools(profile); + self + } + + /// Registers typed application state that tools can retrieve from their context. + pub fn with_context(self, context: Arc) -> Self { + self.handle.register_app_context(context); + self + } + + /// Registers a runtime intrinsic tool. + pub fn with_intrinsic(self, tool: T) -> Self + where + T: ExecutableTool + 'static, + { + self.with_tool(tool) + } + + /// Replaces the runtime store implementation. + /// + /// The default store is not opened on the way here. Recovery runs at build + /// time against whichever store the builder ends with, so a caller that + /// supplies its own never has the machine-wide default database created + /// underneath it. + pub fn with_store(self, store: impl RuntimeStore + 'static) -> Self { + Self { + handle: self.handle.rebind_store(std::sync::Arc::new(store)), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Replaces the command executor used by builtin tools. + pub fn with_executor(self, executor: E) -> Self + where + E: RuntimeExecutor + 'static, + { + Self { + handle: self.handle.with_executor(Arc::new(executor)), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Replaces the compaction engine used for transcript summarization. + pub fn with_compaction_engine(self, engine: C) -> Self + where + C: CompactionEngine + 'static, + { + Self { + handle: self.handle.with_compaction_engine(Arc::new(engine)), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Sets the runtime policy used to authorize file and process access. + pub fn with_policy(self, policy: RuntimePolicy) -> Self { + Self { + handle: self.handle.with_policy(policy), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Installs a pre-tool authorization service for runtime tool calls. + pub fn with_tool_authorizer(self, tool_authorizer: A) -> Self + where + A: ToolAuthorizer + 'static, + { + Self { + handle: self.handle.with_tool_authorizer(Arc::new(tool_authorizer)), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Sets the persisted runtime identifier used to group resumable agents. + pub fn with_runtime_identifier(self, runtime_identifier: impl Into>) -> Self { + Self { + handle: self.handle.with_runtime_identifier(runtime_identifier), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Appends a single runtime hook, keeping any already registered. + pub fn with_hook(self, hook: H) -> Self + where + H: RuntimeHook + 'static, + { + let hooks = self.handle.hooks().clone().with_hook(hook); + Self { + handle: self.handle.with_hooks(hooks), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Appends a single pre-execution hook, keeping any already registered. + pub fn with_pre_hook(self, hook: H) -> Self + where + H: PreExecutionHook + 'static, + { + let pre_hooks = self.handle.pre_hooks().clone().with_hook(hook); + Self { + handle: self.handle.with_pre_hooks(pre_hooks), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Replaces hooks with the provided collection. + pub fn with_hooks(self, hooks: I) -> Self + where + I: IntoIterator>, + { + Self { + handle: self.handle.with_hooks(RuntimeHooks::new().extend(hooks)), + provider_registry: self.provider_registry, + mcp_configs: self.mcp_configs, + } + } + + /// Registers a skills directory and enables the builtin `load_skill` tool. + pub fn with_skills_dir(self, path: impl AsRef) -> Result { + self.handle + .register_skill_loader(SkillLoader::from_dir(path)?); + Ok(self) + } + + /// Registers an MCP server, reached over stdio, to connect to during build. + pub fn with_mcp_server(mut self, config: McpServerConfig) -> Self { + self.mcp_configs + .push(McpRegistration::Stdio(Box::new(config))); + self + } + + /// Registers multiple stdio MCP servers to connect to during build. + pub fn with_mcp_servers(mut self, configs: impl IntoIterator) -> Self { + self.mcp_configs.extend( + configs + .into_iter() + .map(|config| McpRegistration::Stdio(Box::new(config))), + ); + self + } + + /// Registers an MCP server reached over the legacy HTTP+SSE transport. + /// + /// Every tool the server advertises is bridged into the runtime under a + /// namespaced name. Use [`McpSseClient`](crate::mcp::McpSseClient) directly + /// when a host needs to apply its own allowlist before anything is + /// registered. + /// + /// ```rust,no_run + /// use mentra::{BuiltinProvider, McpSseServerConfig, Runtime}; + /// # async fn demo() -> Result<(), Box> { + /// let runtime = Runtime::builder() + /// .with_provider(BuiltinProvider::Anthropic, "sk-...") + /// .with_mcp_sse_server( + /// McpSseServerConfig::new("observability", "https://mcp.example.com/sse") + /// .with_bearer_token(""), + /// ) + /// .build_async() + /// .await?; + /// # let _ = runtime; + /// # Ok(()) + /// # } + /// ``` + pub fn with_mcp_sse_server(mut self, config: McpSseServerConfig) -> Self { + self.mcp_configs + .push(McpRegistration::Sse(Box::new(config))); + self + } + + /// Registers multiple HTTP+SSE MCP servers to connect to during build. + pub fn with_mcp_sse_servers( + mut self, + configs: impl IntoIterator, + ) -> Self { + self.mcp_configs.extend( + configs + .into_iter() + .map(|config| McpRegistration::Sse(Box::new(config))), + ); + self + } + + /// Registers a builtin provider when an API key is present. + pub fn with_optional_provider( + mut self, + id: BuiltinProvider, + api_key: Option>, + ) -> Self { + if let Some(api_key) = api_key { + let _ = self + .provider_registry + .register_builtin_provider(id, api_key.into()); + } + self + } + + /// Registers a builtin provider from an API key. + pub fn with_provider(mut self, id: BuiltinProvider, api_key: impl Into) -> Self { + let _ = self + .provider_registry + .register_builtin_provider(id, api_key); + self + } + + /// Chooses the transport this runtime's Responses-family requests stream + /// over. + /// + /// Runtime scope, because a transport is a property of the connection to a + /// provider rather than of one run: an HTTP+SSE turn and a websocket turn + /// against the same endpoint are two different conversations with it, and a + /// per-run switch would mean the runtime holding two live opinions about + /// one socket. Left unset, each request's own + /// [`ResponsesRequestOptions::transport`](crate::provider::ResponsesRequestOptions) + /// stands — which is HTTP+SSE unless a host set otherwise, exactly as + /// before this method existed. + /// + /// A provider that does not serve websockets — anthropic and gemini, whose + /// definitions report `supports_websockets: false` — refuses an explicit + /// [`ResponsesTransport::WebSocket`](crate::provider::ResponsesTransport) + /// at its first request, naming itself, rather than answering over + /// HTTP+SSE. Selecting a transport is an explicit act, and a silent + /// fallback would hand back a stream nobody asked for. + /// + /// ```rust,no_run + /// use mentra::{BuiltinProvider, Runtime}; + /// use mentra::provider::ResponsesTransport; + /// # fn demo() -> Result<(), Box> { + /// let runtime = Runtime::builder() + /// .with_provider(BuiltinProvider::OpenAI, "sk-...") + /// .with_responses_transport(ResponsesTransport::WebSocket) + /// .build()?; + /// # let _ = runtime; + /// # Ok(()) + /// # } + /// ``` + pub fn with_responses_transport(mut self, transport: ResponsesTransport) -> Self { + self.provider_registry.set_responses_transport(transport); + self + } + + /// Registers the local Ollama provider using its default OpenAI-compatible endpoint. + pub fn with_ollama(mut self) -> Self { + self.provider_registry.register_ollama(); + self + } + + /// Registers the local LM Studio provider using its default OpenAI-compatible endpoint. + pub fn with_lmstudio(mut self) -> Self { + self.provider_registry.register_lmstudio(); + self + } + + /// Registers a custom runtime provider implementation. + /// + /// This is the supported seam for test-time provider injection when you + /// want to script model responses without live API calls. + /// + /// ```rust,no_run + /// use async_trait::async_trait; + /// use mentra::{BuiltinProvider, ModelInfo, ProviderDescriptor, Runtime}; + /// use mentra::error::{ProviderError, RuntimeError}; + /// use mentra::provider::{Provider, ProviderEventStream, Request}; + /// use tokio::sync::mpsc; + /// + /// struct TestProvider; + /// + /// #[async_trait] + /// impl Provider for TestProvider { + /// fn descriptor(&self) -> ProviderDescriptor { + /// ProviderDescriptor::new(BuiltinProvider::Anthropic) + /// } + /// + /// async fn list_models(&self) -> Result, ProviderError> { + /// Ok(vec![ModelInfo::new("test-model", BuiltinProvider::Anthropic)]) + /// } + /// + /// async fn stream( + /// &self, + /// _request: Request<'_>, + /// ) -> Result { + /// let (_tx, rx) = mpsc::unbounded_channel(); + /// Ok(rx) + /// } + /// } + /// + /// let runtime = Runtime::empty_builder() + /// .with_provider_instance(TestProvider) + /// .build()?; + /// # Ok::<(), RuntimeError>(()) + /// ``` + pub fn with_provider_instance

(mut self, provider: P) -> Self + where + P: Provider + 'static, + { + self.provider_registry.register_provider_instance(provider); + self + } + + /// Registers a provider-core instance built from `mentra::provider_core`. + /// + /// Use this when you want Mentra's runtime with a customized provider + /// definition, such as a custom OpenAI-compatible or Anthropic-compatible + /// base URL. + pub fn with_registered_provider

(mut self, provider: P) -> Self + where + P: mentra_provider::Provider + 'static, + { + self.provider_registry + .register_registered_provider(provider); + self + } + + /// Builds the runtime, connects to MCP servers, and validates providers. + /// + /// This is an async method because MCP server connections require spawning + /// processes and performing the initialize handshake. + pub async fn build_async(self) -> Result { + if self.provider_registry.is_empty() { + return Err(RuntimeError::ProviderNotFound(None)); + } + + // Connect to MCP servers and register their tools. + let mut outcomes = Vec::new(); + if !self.mcp_configs.is_empty() { + let mut manager = McpManager::new(); + for config in &self.mcp_configs { + let connected = match config { + McpRegistration::Stdio(config) => manager + .connect(config) + .await + .map_err(|error| error.to_string()), + McpRegistration::Sse(config) => manager + .connect_sse(config) + .await + .map_err(|error| error.to_string()), + }; + + match connected { + Ok(bridged_tools) => { + let tools = bridged_tools.len(); + for tool in bridged_tools { + self.handle.register_tool(tool); + } + outcomes.push(McpServerSummary { + name: config.name().to_string(), + tools, + error: None, + }); + } + Err(error) => { + // Degraded mode: one unreachable server must not sink a + // session. Recorded rather than only printed, so a host + // can say which servers are live instead of a user + // wondering why a tool is missing. + eprintln!( + "Warning: MCP server '{}' failed to connect: {error}", + config.name() + ); + outcomes.push(McpServerSummary { + name: config.name().to_string(), + tools: 0, + error: Some(error), + }); + } + } + } + // Store the manager in the app context for later use. + self.handle + .register_app_context(Arc::new(tokio::sync::Mutex::new(manager))); + } + + let provider_registry = Arc::new(std::sync::RwLock::new(self.provider_registry)); + let handle = self + .handle + .with_provider_registry(provider_registry.clone()); + handle.prepare_recovery(); + Ok(Runtime { + handle, + provider_registry, + mcp_servers: outcomes, + }) + } + + /// Builds the runtime synchronously. + /// + /// Connecting to an MCP server means spawning a process and completing a + /// handshake, which cannot happen here — so registering one and then + /// calling this is refused rather than silently honored halfway. Use + /// [`build_async`](Self::build_async) when MCP servers are configured. + pub fn build(self) -> Result { + if self.provider_registry.is_empty() { + return Err(RuntimeError::ProviderNotFound(None)); + } + + if !self.mcp_configs.is_empty() { + let names: Vec<&str> = self.mcp_configs.iter().map(McpRegistration::name).collect(); + return Err(RuntimeError::OperationDenied(format!( + "MCP servers are registered ({}) but `build` cannot connect them; \ + use `build_async`", + names.join(", ") + ))); + } + + let provider_registry = Arc::new(std::sync::RwLock::new(self.provider_registry)); + let handle = self + .handle + .with_provider_registry(provider_registry.clone()); + handle.prepare_recovery(); + Ok(Runtime { + handle, + provider_registry, + mcp_servers: Vec::new(), + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::runtime::VolatileRuntimeStore; + use crate::runtime::control::{HookDecision, PreExecutionContext}; + use crate::runtime::store::default_store_paths_on_this_thread; + use async_trait::async_trait; + use std::sync::atomic::{AtomicUsize, Ordering}; + + /// The least a builder will accept: a provider must exist before any + /// other check runs. + struct StubProvider; + + #[async_trait] + impl crate::provider::Provider for StubProvider { + fn descriptor(&self) -> crate::provider::ProviderDescriptor { + crate::provider::ProviderDescriptor::new(BuiltinProvider::OpenAI) + } + + async fn list_models( + &self, + ) -> Result, crate::provider::ProviderError> { + Ok(Vec::new()) + } + + async fn stream( + &self, + _request: crate::provider::Request<'_>, + ) -> Result { + unreachable!("no turn is run in these tests") + } + } + + /// Counts how many times it was consulted, so a hook that was silently + /// dropped during registration shows up as a count that never moves. + struct Counting(Arc); + + #[async_trait] + impl PreExecutionHook for Counting { + async fn pre_tool_execution( + &self, + _context: &PreExecutionContext, + ) -> Result { + self.0.fetch_add(1, Ordering::SeqCst); + Ok(HookDecision::Allow) + } + } + + #[tokio::test] + async fn registering_a_second_pre_hook_keeps_the_first() { + let first = Arc::new(AtomicUsize::new(0)); + let second = Arc::new(AtomicUsize::new(0)); + + let builder = RuntimeBuilder::new(false) + .with_pre_hook(Counting(Arc::clone(&first))) + .with_pre_hook(Counting(Arc::clone(&second))); + + let context = PreExecutionContext { + agent_id: "a1".to_string(), + tool_name: "shell".to_string(), + tool_call_id: "tc-1".to_string(), + input_json: "{}".to_string(), + working_directory: std::path::PathBuf::from("/repo"), + }; + builder + .handle + .pre_hooks() + .run(&context) + .await + .expect("hooks run"); + + // The first registration used to be discarded by the second, which is + // a security-relevant silent failure for a veto seam. + assert_eq!( + first.load(Ordering::SeqCst), + 1, + "the first hook must still run" + ); + assert_eq!(second.load(Ordering::SeqCst), 1); + } + + #[test] + fn build_refuses_to_discard_registered_mcp_servers() { + let error = RuntimeBuilder::new(false) + .with_provider_instance(StubProvider) + .with_mcp_server(McpServerConfig { + name: "github".to_string(), + command: "npx".to_string(), + args: Vec::new(), + env: Default::default(), + cwd: None, + }) + .build() + .err() + .expect("a sync build cannot connect a server, so it must say so"); + + // The old behavior was to build cleanly and drop the server, which a + // caller only discovered when a tool it had configured was missing. + let message = error.to_string(); + assert!( + message.contains("github") && message.contains("build_async"), + "the refusal must name the server and the way forward: {message}" + ); + } + + #[tokio::test] + async fn a_runtime_with_no_mcp_servers_reports_none() { + let runtime = RuntimeBuilder::new(false) + .with_provider_instance(StubProvider) + .build_async() + .await + .expect("builds"); + + assert!(runtime.mcp_servers().is_empty()); + } + + /// A caller that supplies its own store has opted out of the machine-wide + /// default. Constructing the handle used to open it anyway — creating + /// `runtime.sqlite` on a pristine machine — before `with_store` replaced + /// the store it had just prepared. + #[test] + fn a_build_with_a_caller_store_leaves_the_default_database_alone() { + let store = VolatileRuntimeStore::new(); + let probe = store.clone(); + + let runtime = RuntimeBuilder::new(false) + .with_store(store) + .with_provider_instance(StubProvider) + .build() + .expect("builds"); + + let default_paths = default_store_paths_on_this_thread(); + assert!( + !default_paths.is_empty(), + "the handle still constructs a default store, so this test has something to check" + ); + for path in default_paths { + assert!( + !path.exists(), + "a discarded default store must never be opened: {}", + path.display() + ); + } + assert_eq!( + probe.recovery_preparations(), + 1, + "recovery must run once, on the store the caller kept" + ); + drop(runtime); + } + + /// The async build boundary carries the same guarantee as the sync one. + #[tokio::test] + async fn an_async_build_prepares_recovery_once_on_the_caller_store() { + let store = VolatileRuntimeStore::new(); + let probe = store.clone(); + + let runtime = RuntimeBuilder::new(false) + .with_store(store) + .with_provider_instance(StubProvider) + .build_async() + .await + .expect("builds"); + + assert_eq!(probe.recovery_preparations(), 1); + drop(runtime); + } + + /// Deferring recovery must not skip it: a build that keeps the default + /// store still reconciles interrupted state, which for SQLite means the + /// database is opened and its schema created. + #[test] + fn a_default_build_still_prepares_recovery() { + let runtime = RuntimeBuilder::new(false) + .with_provider_instance(StubProvider) + .build() + .expect("builds"); + + let default_paths = default_store_paths_on_this_thread(); + assert!(!default_paths.is_empty()); + for path in default_paths { + assert!( + path.exists(), + "the store a runtime actually kept must be prepared: {}", + path.display() + ); + } + drop(runtime); + } + + /// Recovery belongs to the build boundary, not to assembly: until `build` + /// settles which store survives, nothing may be prepared. + #[test] + fn assembling_a_builder_prepares_nothing() { + let store = VolatileRuntimeStore::new(); + let probe = store.clone(); + + let builder = RuntimeBuilder::new(false) + .with_store(store) + .with_provider_instance(StubProvider); + + assert_eq!( + probe.recovery_preparations(), + 0, + "a store is only prepared once the builder is done being reconfigured" + ); + drop(builder); + } + + /// A store that is swapped out again must never be prepared: `with_store` + /// used to prepare eagerly, which made every intermediate store pay for a + /// choice the caller went on to revise. + #[test] + fn a_replaced_store_is_never_prepared() { + let discarded = VolatileRuntimeStore::new(); + let discarded_probe = discarded.clone(); + let kept = VolatileRuntimeStore::new(); + let kept_probe = kept.clone(); + + let runtime = RuntimeBuilder::new(false) + .with_store(discarded) + .with_store(kept) + .with_provider_instance(StubProvider) + .build() + .expect("builds"); + + assert_eq!( + discarded_probe.recovery_preparations(), + 0, + "a store the builder threw away must not have been opened" + ); + assert_eq!(kept_probe.recovery_preparations(), 1); + drop(runtime); + } +} diff --git a/vendor/mentra/src/runtime/control.rs b/vendor/mentra/src/runtime/control.rs new file mode 100644 index 0000000..dafd33e --- /dev/null +++ b/vendor/mentra/src/runtime/control.rs @@ -0,0 +1,19 @@ +mod command; +mod hooks; +mod policy; +mod run; +/// Container and sandbox environment detection. +pub mod sandbox; + +pub use command::{ + CommandOutput, CommandRequest, CommandSpec, ExecOutput, LocalRuntimeExecutor, RuntimeExecutor, + read_limited_file, +}; +pub use hooks::{ + AuditHook, AuditLogHook, HookDecision, PreExecutionContext, PreExecutionHook, + PreExecutionHooks, RuntimeHook, RuntimeHookEvent, RuntimeHooks, is_transient_provider_error, + is_transient_runtime_error, +}; +pub(crate) use policy::ShellValidation; +pub use policy::{RuntimePolicy, ShellValidationMode}; +pub use run::{CancellationFlag, CancellationToken, EarlyEnd, ProviderRetry, RunOptions}; diff --git a/vendor/mentra/src/runtime/control/command.rs b/vendor/mentra/src/runtime/control/command.rs new file mode 100644 index 0000000..81dfae1 --- /dev/null +++ b/vendor/mentra/src/runtime/control/command.rs @@ -0,0 +1,476 @@ +use std::{ + io, + path::{Path, PathBuf}, + time::Duration, +}; + +#[cfg(windows)] +use std::process::Command as StdCommand; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use tokio::{ + io::{AsyncBufReadExt, AsyncRead, AsyncReadExt}, + process::{Child, Command}, +}; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ExecOutput { + pub stdout: String, + pub stderr: String, + pub success: bool, + pub status_code: Option, + pub timed_out: bool, + pub stdout_truncated: bool, + pub stderr_truncated: bool, +} + +impl ExecOutput { + pub fn success(&self) -> bool { + self.success + } +} + +pub type CommandOutput = ExecOutput; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum CommandSpec { + Shell { command: String }, +} + +impl CommandSpec { + pub fn display(&self) -> &str { + match self { + Self::Shell { command } => command, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct CommandRequest { + pub spec: CommandSpec, + pub cwd: PathBuf, + pub timeout: Duration, + pub env: Vec<(String, String)>, + pub max_output_bytes_per_stream: usize, + /// Where the host asked this command to run; `None` is the local executor. + /// + /// Execution data, not policy: the executor reads it, nothing else decides + /// on it. A targeted request is authorized, validated, timeout-clamped and + /// output-capped exactly like a local one, so routing a command elsewhere + /// can never be a way around the policy that guards running it here. An + /// executor that does not serve the named target must refuse the request + /// rather than run it locally. + /// + /// Defaulted on deserialization so a request serialized before this field + /// existed still loads, as the untargeted request it was. + #[serde(default)] + pub target: Option, +} + +/// Executes runtime command requests. +/// +/// Implementations are trusted host components. A sandboxed implementation +/// should be configured with an immutable filesystem and network policy because +/// [`CommandRequest`] intentionally carries execution data, not authorization +/// policy. +#[async_trait] +pub trait RuntimeExecutor: Send + Sync { + async fn run(&self, request: CommandRequest) -> Result; + + /// Runs an untargeted command. + /// + /// The convenience form keeps the signature it always had, so it can only + /// build a request with [`CommandRequest::target`] set to `None`. A caller + /// that needs a target builds the [`CommandRequest`] itself and calls + /// [`run`](Self::run). + async fn run_command( + &self, + command: &str, + cwd: &Path, + timeout: Duration, + env: Vec<(String, String)>, + max_output_bytes_per_stream: usize, + ) -> Result { + self.run(CommandRequest { + spec: CommandSpec::Shell { + command: command.to_string(), + }, + cwd: cwd.to_path_buf(), + timeout, + env, + max_output_bytes_per_stream, + target: None, + }) + .await + } +} + +/// Executes commands directly with the current user's host permissions. +/// +/// This executor clears unlisted environment variables and enforces output, +/// timeout, and timeout-cleanup limits. It does not sandbox filesystem or +/// network access. +/// +/// It serves no named target and refuses any request that carries one: a +/// command the host addressed elsewhere silently running on this machine +/// would be the one failure mode a target is meant to prevent. +pub struct LocalRuntimeExecutor; + +#[async_trait] +impl RuntimeExecutor for LocalRuntimeExecutor { + async fn run(&self, request: CommandRequest) -> Result { + let CommandRequest { + spec, + cwd, + timeout, + env, + max_output_bytes_per_stream, + target, + } = request; + if let Some(target) = target { + return Err(format!( + "no executor serves target `{target}`; the local executor only runs untargeted commands" + )); + } + let command = match spec { + CommandSpec::Shell { command } => command, + }; + + let mut process = Command::new(platform_shell_program()); + process + .args(platform_shell_args(&command)) + .current_dir(&cwd) + .env_clear() + .envs(env) + .stdin(std::process::Stdio::null()) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()) + .kill_on_drop(true); + + #[cfg(unix)] + { + unsafe { + process.pre_exec(|| { + if libc::setsid() == -1 { + return Err(io::Error::last_os_error()); + } + Ok(()) + }); + } + } + + let mut child = process + .spawn() + .map_err(|error| format!("Failed to execute command: {error}"))?; + + let stdout = child + .stdout + .take() + .ok_or_else(|| "Failed to capture stdout".to_string())?; + let stderr = child + .stderr + .take() + .ok_or_else(|| "Failed to capture stderr".to_string())?; + let stdout_task = tokio::spawn(read_capped(stdout, max_output_bytes_per_stream)); + let stderr_task = tokio::spawn(read_capped(stderr, max_output_bytes_per_stream)); + + let wait_result = tokio::time::timeout(timeout, child.wait()).await; + let timed_out = wait_result.is_err(); + let status = if timed_out { + kill_entire_process_tree(&mut child) + .map_err(|error| format!("Failed to stop timed out command: {error}"))?; + let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await; + None + } else { + Some( + wait_result + .expect("non-timeout wait result") + .map_err(|error| format!("Failed to wait for command: {error}"))?, + ) + }; + + let stdout = join_stream(stdout_task).await?; + let stderr = join_stream(stderr_task).await?; + + let (success, status_code) = if timed_out { + (false, Some(124)) + } else if let Some(status) = status { + (status.success(), status.code()) + } else { + (false, None) + }; + + Ok(CommandOutput { + stdout: String::from_utf8_lossy(&stdout.bytes).into_owned(), + stderr: String::from_utf8_lossy(&stderr.bytes).into_owned(), + success, + status_code, + timed_out, + stdout_truncated: stdout.truncated, + stderr_truncated: stderr.truncated, + }) + } +} + +struct StreamCapture { + bytes: Vec, + truncated: bool, +} + +async fn read_capped(mut reader: R, max_bytes: usize) -> io::Result +where + R: AsyncRead + Unpin + Send + 'static, +{ + let mut bytes = Vec::new(); + let mut truncated = false; + let mut buffer = [0u8; 8192]; + + loop { + let read = reader.read(&mut buffer).await?; + if read == 0 { + break; + } + + let remaining = max_bytes.saturating_sub(bytes.len()); + let take = remaining.min(read); + bytes.extend_from_slice(&buffer[..take]); + if take < read { + truncated = true; + } + } + + Ok(StreamCapture { bytes, truncated }) +} + +async fn join_stream( + handle: tokio::task::JoinHandle>, +) -> Result { + tokio::time::timeout(Duration::from_secs(2), handle) + .await + .map_err(|_| "Timed out while draining command output".to_string())? + .map_err(|error| format!("Failed to join command output task: {error}"))? + .map_err(|error| format!("Failed to read command output: {error}")) +} + +fn kill_entire_process_tree(child: &mut Child) -> io::Result<()> { + #[cfg(unix)] + { + if let Some(pid) = child.id() { + let result = unsafe { libc::kill(-(pid as i32), libc::SIGKILL) }; + if result == -1 { + let error = io::Error::last_os_error(); + if error.raw_os_error() != Some(libc::ESRCH) { + return Err(error); + } + } + } + } + + #[cfg(windows)] + { + if let Some(pid) = child.id() { + let status = StdCommand::new("taskkill") + .args(["/PID", &pid.to_string(), "/T", "/F"]) + .status()?; + if status.success() { + return Ok(()); + } + + if child.try_wait()?.is_some() { + return Ok(()); + } + } + } + + child.start_kill() +} + +#[cfg(unix)] +fn platform_shell_program() -> &'static str { + "/bin/sh" +} + +#[cfg(windows)] +fn platform_shell_program() -> &'static str { + "cmd.exe" +} + +#[cfg(unix)] +fn platform_shell_args(command: &str) -> [&str; 2] { + ["-c", command] +} + +#[cfg(windows)] +fn platform_shell_args(command: &str) -> [&str; 2] { + ["/C", command] +} + +pub async fn read_limited_file(path: &Path, max_lines: Option) -> Result { + let file = tokio::fs::File::open(path) + .await + .map_err(|error| format!("Failed to open file: {error}"))?; + let mut lines = tokio::io::BufReader::new(file).lines(); + let mut content = Vec::new(); + + loop { + if let Some(limit) = max_lines + && content.len() >= limit + { + break; + } + + match lines.next_line().await { + Ok(Some(line)) => content.push(line), + Ok(None) => break, + Err(error) => return Err(format!("Failed to read file: {error}")), + } + } + + Ok(content.join("\n")) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[cfg(unix)] + fn stdout_and_stderr_command() -> String { + "printf 'aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa'; printf 'bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb' >&2" + .to_string() + } + + #[cfg(windows)] + fn stdout_and_stderr_command() -> String { + "echo aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa& echo bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb 1>&2" + .to_string() + } + + #[cfg(unix)] + fn missing_secret_command() -> String { + "printf '%s' \"${SECRET:-missing}\"".to_string() + } + + #[cfg(windows)] + fn missing_secret_command() -> String { + "if defined SECRET (echo unexpected) else (echo missing)".to_string() + } + + #[cfg(unix)] + fn timeout_command() -> String { + "sleep 1".to_string() + } + + #[cfg(windows)] + fn timeout_command() -> String { + "ping.exe -n 2 127.0.0.1 >nul".to_string() + } + + #[cfg(unix)] + fn minimal_shell_env() -> Vec<(String, String)> { + vec![( + "PATH".to_string(), + std::env::var("PATH").expect("path available"), + )] + } + + #[cfg(windows)] + fn minimal_shell_env() -> Vec<(String, String)> { + ["PATH", "PATHEXT", "SystemRoot", "COMSPEC", "TEMP", "TMP"] + .into_iter() + .filter_map(|name| { + std::env::var(name) + .ok() + .map(|value| (name.to_string(), value)) + }) + .collect() + } + + #[tokio::test] + async fn caps_stdout_and_stderr_independently() { + let output = LocalRuntimeExecutor + .run(CommandRequest { + spec: CommandSpec::Shell { + command: stdout_and_stderr_command(), + }, + cwd: std::env::temp_dir(), + timeout: Duration::from_secs(5), + env: minimal_shell_env(), + max_output_bytes_per_stream: 8, + target: None, + }) + .await + .expect("command output"); + + assert!(!output.timed_out, "{output:?}"); + assert!(output.success, "{output:?}"); + assert_eq!(output.stdout.len(), 8); + assert_eq!(output.stderr.len(), 8); + assert!(output.stdout_truncated); + assert!(output.stderr_truncated); + } + + #[tokio::test] + async fn allowlisted_environment_is_enforced() { + let output = LocalRuntimeExecutor + .run(CommandRequest { + spec: CommandSpec::Shell { + command: missing_secret_command(), + }, + cwd: std::env::temp_dir(), + timeout: Duration::from_secs(5), + env: minimal_shell_env(), + max_output_bytes_per_stream: 1024, + target: None, + }) + .await + .expect("command output"); + + assert!(!output.timed_out, "{output:?}"); + assert!(output.success, "{output:?}"); + assert_eq!(output.stdout.trim_end(), "missing"); + } + + #[tokio::test] + async fn timeout_marks_output_and_uses_timeout_exit_code() { + let output = LocalRuntimeExecutor + .run(CommandRequest { + spec: CommandSpec::Shell { + command: timeout_command(), + }, + cwd: std::env::temp_dir(), + timeout: Duration::from_millis(50), + env: minimal_shell_env(), + max_output_bytes_per_stream: 1024, + target: None, + }) + .await + .expect("command output"); + + assert!(output.timed_out); + assert_eq!(output.status_code, Some(124)); + assert!(!output.success); + } + + #[tokio::test] + async fn targeted_request_is_refused_instead_of_running_locally() { + let error = LocalRuntimeExecutor + .run(CommandRequest { + spec: CommandSpec::Shell { + command: "printf 'ran locally'".to_string(), + }, + cwd: std::env::temp_dir(), + timeout: Duration::from_secs(5), + env: minimal_shell_env(), + max_output_bytes_per_stream: 1024, + target: Some("mac".to_string()), + }) + .await + .expect_err("a targeted request must not run locally"); + + assert_eq!( + error, + "no executor serves target `mac`; the local executor only runs untargeted commands" + ); + } +} diff --git a/vendor/mentra/src/runtime/control/hooks.rs b/vendor/mentra/src/runtime/control/hooks.rs new file mode 100644 index 0000000..3d94193 --- /dev/null +++ b/vendor/mentra/src/runtime/control/hooks.rs @@ -0,0 +1,696 @@ +use std::{path::PathBuf, sync::Arc}; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; + +use crate::{ + provider::{ProviderError, TokenUsage}, + runtime::{AuditStore, RuntimeStore, error::RuntimeError}, + tool::{ToolAuthorizationOutcome, ToolAuthorizationPreview}, +}; + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "kind", rename_all = "snake_case")] +pub enum RuntimeHookEvent { + AuthorizationDenied { + agent_id: String, + action: String, + detail: String, + }, + ToolAuthorizationStarted { + agent_id: String, + tool_name: String, + tool_call_id: String, + preview: ToolAuthorizationPreview, + }, + ToolAuthorizationFinished { + agent_id: String, + tool_name: String, + tool_call_id: String, + outcome: ToolAuthorizationOutcome, + reason: Option, + }, + ToolAuthorizationBlocked { + agent_id: String, + tool_name: String, + tool_call_id: String, + outcome: ToolAuthorizationOutcome, + reason: Option, + }, + RecoveryPrepared { + runtime_instance_id: String, + }, + ModelRequestStarted { + agent_id: String, + model: String, + attempt: usize, + }, + ModelRequestFinished { + agent_id: String, + model: String, + attempt: usize, + success: bool, + error: Option, + }, + ModelResponseFinished { + agent_id: String, + model: String, + attempt: usize, + success: bool, + error: Option, + stop_reason: Option, + usage: Option, + }, + ToolExecutionStarted { + agent_id: String, + tool_name: String, + tool_call_id: String, + }, + ToolExecutionFinished { + agent_id: String, + tool_name: String, + tool_call_id: String, + is_error: bool, + error: Option, + output_preview: String, + /// Opaque host metadata attached via `ToolOutput::details`, carried up + /// to this observability boundary. Never sent to a provider. + #[serde(default)] + details: Option, + }, + PolicyDenied { + agent_id: String, + tool_name: String, + reason: String, + }, + BackgroundTaskStarted { + agent_id: String, + task_id: String, + command: String, + cwd: PathBuf, + }, + BackgroundTaskFinished { + agent_id: String, + task_id: String, + status: String, + }, + MemorySearchStarted { + agent_id: String, + limit: usize, + query_preview: String, + }, + MemorySearchFinished { + agent_id: String, + success: bool, + result_count: usize, + error: Option, + }, + MemoryIngestStarted { + agent_id: String, + source_revision: u64, + }, + MemoryIngestFinished { + agent_id: String, + source_revision: u64, + success: bool, + stored_records: usize, + error: Option, + }, + MemoryCompactionProposed { + agent_id: String, + base_revision: u64, + transcript_path: PathBuf, + }, + MemoryCompactionApplied { + agent_id: String, + base_revision: u64, + resulting_history_len: usize, + }, + MemoryCompactionSkipped { + agent_id: String, + base_revision: u64, + }, + RunAborted { + agent_id: String, + reason: String, + }, + ToolExecutionBlocked { + agent_id: String, + tool_name: String, + tool_call_id: String, + reason: String, + }, +} + +impl RuntimeHookEvent { + fn scope(&self) -> String { + match self { + Self::AuthorizationDenied { agent_id, .. } => agent_id.clone(), + Self::ToolAuthorizationStarted { agent_id, .. } => agent_id.clone(), + Self::ToolAuthorizationFinished { agent_id, .. } => agent_id.clone(), + Self::ToolAuthorizationBlocked { agent_id, .. } => agent_id.clone(), + Self::RecoveryPrepared { + runtime_instance_id, + } => runtime_instance_id.clone(), + Self::ModelRequestStarted { agent_id, .. } + | Self::ModelRequestFinished { agent_id, .. } + | Self::ModelResponseFinished { agent_id, .. } + | Self::ToolExecutionStarted { agent_id, .. } + | Self::ToolExecutionFinished { agent_id, .. } + | Self::PolicyDenied { agent_id, .. } + | Self::BackgroundTaskStarted { agent_id, .. } + | Self::BackgroundTaskFinished { agent_id, .. } + | Self::MemorySearchStarted { agent_id, .. } + | Self::MemorySearchFinished { agent_id, .. } + | Self::MemoryIngestStarted { agent_id, .. } + | Self::MemoryIngestFinished { agent_id, .. } + | Self::MemoryCompactionProposed { agent_id, .. } + | Self::MemoryCompactionApplied { agent_id, .. } + | Self::MemoryCompactionSkipped { agent_id, .. } + | Self::RunAborted { agent_id, .. } + | Self::ToolExecutionBlocked { agent_id, .. } => agent_id.clone(), + } + } + + fn event_type(&self) -> &'static str { + match self { + Self::AuthorizationDenied { .. } => "authorization_denied", + Self::ToolAuthorizationStarted { .. } => "tool_authorization_started", + Self::ToolAuthorizationFinished { .. } => "tool_authorization_finished", + Self::ToolAuthorizationBlocked { .. } => "tool_authorization_blocked", + Self::RecoveryPrepared { .. } => "recovery_prepared", + Self::ModelRequestStarted { .. } => "model_request_started", + Self::ModelRequestFinished { .. } => "model_request_finished", + Self::ModelResponseFinished { .. } => "model_response_finished", + Self::ToolExecutionStarted { .. } => "tool_execution_started", + Self::ToolExecutionFinished { .. } => "tool_execution_finished", + Self::PolicyDenied { .. } => "policy_denied", + Self::BackgroundTaskStarted { .. } => "background_task_started", + Self::BackgroundTaskFinished { .. } => "background_task_finished", + Self::MemorySearchStarted { .. } => "memory_search_started", + Self::MemorySearchFinished { .. } => "memory_search_finished", + Self::MemoryIngestStarted { .. } => "memory_ingest_started", + Self::MemoryIngestFinished { .. } => "memory_ingest_finished", + Self::MemoryCompactionProposed { .. } => "memory_compaction_proposed", + Self::MemoryCompactionApplied { .. } => "memory_compaction_applied", + Self::MemoryCompactionSkipped { .. } => "memory_compaction_skipped", + Self::RunAborted { .. } => "run_aborted", + Self::ToolExecutionBlocked { .. } => "tool_execution_blocked", + } + } +} + +pub trait RuntimeHook: Send + Sync { + fn on_event( + &self, + store: &dyn AuditStore, + event: &RuntimeHookEvent, + ) -> Result<(), RuntimeError>; +} + +pub struct AuditHook; +pub type AuditLogHook = AuditHook; + +impl RuntimeHook for AuditHook { + fn on_event( + &self, + store: &dyn AuditStore, + event: &RuntimeHookEvent, + ) -> Result<(), RuntimeError> { + store.record_audit_event( + &event.scope(), + event.event_type(), + serde_json::to_value(event).map_err(|error| RuntimeError::Store(error.to_string()))?, + ) + } +} + +#[derive(Clone, Default)] +pub struct RuntimeHooks { + hooks: Vec>, +} + +impl RuntimeHooks { + pub fn new() -> Self { + Self { hooks: Vec::new() } + } + + pub fn with_hook(mut self, hook: H) -> Self + where + H: RuntimeHook + 'static, + { + self.hooks.push(Arc::new(hook)); + self + } + + pub fn extend(mut self, hooks: I) -> Self + where + I: IntoIterator>, + { + self.hooks.extend(hooks); + self + } + + pub fn emit( + &self, + store: &dyn AuditStore, + event: &RuntimeHookEvent, + ) -> Result<(), RuntimeError> { + for hook in &self.hooks { + hook.on_event(store, event)?; + } + Ok(()) + } + + pub(crate) fn emit_runtime( + &self, + store: &dyn RuntimeStore, + event: &RuntimeHookEvent, + ) -> Result<(), RuntimeError> { + self.emit(&RuntimeAuditStore(store), event) + } +} + +/// Stable adapter for Rust versions that cannot upcast `dyn RuntimeStore` to +/// its `AuditStore` supertrait object directly. +struct RuntimeAuditStore<'a>(&'a dyn RuntimeStore); + +impl AuditStore for RuntimeAuditStore<'_> { + fn record_audit_event( + &self, + scope: &str, + event_type: &str, + payload: serde_json::Value, + ) -> Result<(), RuntimeError> { + self.0.record_audit_event(scope, event_type, payload) + } +} + +// --------------------------------------------------------------------------- +// Pre-execution hook types +// --------------------------------------------------------------------------- + +#[derive(Debug, Clone)] +pub struct PreExecutionContext { + pub agent_id: String, + pub tool_name: String, + pub tool_call_id: String, + pub input_json: String, + /// What a relative path in `input_json` resolves against. + /// + /// A hook inspecting `{"path": "../../etc/hosts"}` cannot judge it without + /// knowing where it starts from, and guessing the workspace root is only + /// right until it isn't. + pub working_directory: PathBuf, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum HookDecision { + Allow, + Deny(String), + /// Run the tool with this input instead of the one the model produced. + /// + /// For the cases a veto answers badly: redacting a secret out of an + /// argument, normalizing a path against the right root, narrowing an + /// over-broad command. Denying those costs a round trip and often does not + /// converge, because the model is told "no" without being told what would + /// have been acceptable. + /// + /// The replacement is re-checked by every remaining hook, so a later hook + /// still sees — and can still refuse — what an earlier one produced. A hook + /// cannot use `Modify` to smuggle a call past a hook that runs after it. + Modify { + /// The tool's new input, as JSON. + input_json: String, + /// Why, for the audit trail. + reason: Option, + }, +} + +/// Consulted before a tool runs, and able to stop or rewrite the call. +/// +/// Async because it is invoked from inside a turn: a hook that reads a file, +/// spawns a process, or asks a service would otherwise block a runtime worker +/// for its whole duration. A synchronous signature left every implementor to +/// discover that `tokio::task::block_in_place` panics on a current_thread +/// runtime and to branch on `Handle::runtime_flavor()` themselves. +/// +/// The same shape as [`ToolAuthorizer`](crate::tool::ToolAuthorizer), which +/// sits at the adjacent seam doing the same kind of work. +#[async_trait] +pub trait PreExecutionHook: Send + Sync { + async fn pre_tool_execution( + &self, + context: &PreExecutionContext, + ) -> Result; +} + +/// Forwards to the hook inside. +/// +/// Lets a caller hold a hook it chose at runtime — one of several, or none — +/// and still hand it to anything taking `impl PreExecutionHook`, without each +/// caller writing this impl itself. The same courtesy `ToolAuthorizer` gets. +#[async_trait] +impl PreExecutionHook for Box { + async fn pre_tool_execution( + &self, + context: &PreExecutionContext, + ) -> Result { + (**self).pre_tool_execution(context).await + } +} + +#[async_trait] +impl PreExecutionHook for Arc { + async fn pre_tool_execution( + &self, + context: &PreExecutionContext, + ) -> Result { + (**self).pre_tool_execution(context).await + } +} + +#[derive(Clone, Default)] +pub struct PreExecutionHooks { + hooks: Vec>, +} + +impl PreExecutionHooks { + pub fn new() -> Self { + Self { hooks: Vec::new() } + } + + pub fn with_hook(mut self, hook: H) -> Self + where + H: PreExecutionHook + 'static, + { + self.hooks.push(Arc::new(hook)); + self + } + + /// Runs every hook in order, threading any modification through the rest. + /// + /// Returns the surviving decision: a `Deny` from any hook short-circuits, + /// and otherwise the last `Modify` (if any) is what the tool should run + /// with. Each hook sees the input as its predecessors left it, so + /// modifications compose and no hook can route a call around a later one. + pub async fn run(&self, context: &PreExecutionContext) -> Result { + let mut current = context.clone(); + let mut modified = None; + + for hook in &self.hooks { + match hook.pre_tool_execution(¤t).await? { + HookDecision::Allow => continue, + deny @ HookDecision::Deny(_) => return Ok(deny), + HookDecision::Modify { input_json, reason } => { + current.input_json = input_json.clone(); + modified = Some(HookDecision::Modify { input_json, reason }); + } + } + } + + Ok(modified.unwrap_or(HookDecision::Allow)) + } + + #[allow(dead_code)] + pub fn is_empty(&self) -> bool { + self.hooks.is_empty() + } +} + +/// Returns whether a provider error is likely transient and worth retrying. +pub fn is_transient_provider_error(error: &ProviderError) -> bool { + match error { + ProviderError::Transport(_) + | ProviderError::Decode(_) + | ProviderError::Retryable { .. } => true, + ProviderError::Http { status, .. } => { + status.is_server_error() + || *status == reqwest::StatusCode::TOO_MANY_REQUESTS + || *status == reqwest::StatusCode::REQUEST_TIMEOUT + } + ProviderError::Serialize(_) + | ProviderError::Deserialize(_) + | ProviderError::InvalidRequest(_) + | ProviderError::InvalidResponse(_) + | ProviderError::MalformedStream(_) + | ProviderError::UnsupportedCapability(_) => false, + } +} + +/// Returns whether a runtime error is backed by a transient provider failure. +/// +/// Delegates to [`RuntimeError::category()`] so there is a single source of +/// truth for error classification. +pub fn is_transient_runtime_error(error: &RuntimeError) -> bool { + error.category() == crate::error::ErrorCategory::Retryable +} + +#[cfg(test)] +mod tests { + use super::*; + use std::time::Duration; + + fn make_context(tool_name: &str) -> PreExecutionContext { + PreExecutionContext { + agent_id: "agent-1".to_string(), + tool_name: tool_name.to_string(), + tool_call_id: "call-1".to_string(), + input_json: "{}".to_string(), + working_directory: PathBuf::from("/repo"), + } + } + + struct AllowHook; + #[async_trait] + impl PreExecutionHook for AllowHook { + async fn pre_tool_execution( + &self, + _context: &PreExecutionContext, + ) -> Result { + Ok(HookDecision::Allow) + } + } + + struct DenyHook; + #[async_trait] + impl PreExecutionHook for DenyHook { + async fn pre_tool_execution( + &self, + _context: &PreExecutionContext, + ) -> Result { + Ok(HookDecision::Deny("denied by DenyHook".to_string())) + } + } + + struct ToolNameDenyHook { + blocked_tool: String, + } + #[async_trait] + impl PreExecutionHook for ToolNameDenyHook { + async fn pre_tool_execution( + &self, + context: &PreExecutionContext, + ) -> Result { + if context.tool_name == self.blocked_tool { + Ok(HookDecision::Deny(format!( + "tool '{}' is blocked", + context.tool_name + ))) + } else { + Ok(HookDecision::Allow) + } + } + } + + #[tokio::test] + async fn empty_pre_hooks_allows() { + let hooks = PreExecutionHooks::new(); + let result = hooks.run(&make_context("shell")).await.unwrap(); + assert_eq!(result, HookDecision::Allow); + } + + #[tokio::test] + async fn all_allow_hooks_allows() { + let hooks = PreExecutionHooks::new() + .with_hook(AllowHook) + .with_hook(AllowHook); + let result = hooks.run(&make_context("files")).await.unwrap(); + assert_eq!(result, HookDecision::Allow); + } + + #[tokio::test] + async fn first_deny_wins() { + let hooks = PreExecutionHooks::new() + .with_hook(AllowHook) + .with_hook(DenyHook) + .with_hook(AllowHook); + let result = hooks.run(&make_context("any_tool")).await.unwrap(); + assert_eq!(result, HookDecision::Deny("denied by DenyHook".to_string())); + } + + #[tokio::test] + async fn conditional_deny_by_tool_name() { + let hooks = PreExecutionHooks::new().with_hook(ToolNameDenyHook { + blocked_tool: "shell".to_string(), + }); + + let shell_result = hooks.run(&make_context("shell")).await.unwrap(); + assert_eq!( + shell_result, + HookDecision::Deny("tool 'shell' is blocked".to_string()) + ); + + let files_result = hooks.run(&make_context("files")).await.unwrap(); + assert_eq!(files_result, HookDecision::Allow); + } + + fn http(status: reqwest::StatusCode, retry_after: Option) -> ProviderError { + ProviderError::Http { + status, + body: String::new(), + retry_after, + } + } + + #[test] + fn a_rate_limit_is_transient_whether_or_not_it_named_a_window() { + // Classification is what decides a retry happens at all; the schedule + // only decides how long it waits. A `Retry-After` must not change the + // first answer, in either direction. + for retry_after in [None, Some(Duration::from_secs(45))] { + assert!(is_transient_provider_error(&http( + reqwest::StatusCode::TOO_MANY_REQUESTS, + retry_after + ))); + assert!(is_transient_provider_error(&http( + reqwest::StatusCode::SERVICE_UNAVAILABLE, + retry_after + ))); + } + + assert!( + !is_transient_provider_error(&http(reqwest::StatusCode::BAD_REQUEST, None)), + "a request the caller must fix is not worth re-sending" + ); + } +} + +#[cfg(test)] +mod pre_execution_tests { + use super::*; + + struct Fixed(HookDecision); + + #[async_trait] + impl PreExecutionHook for Fixed { + async fn pre_tool_execution( + &self, + _context: &PreExecutionContext, + ) -> Result { + Ok(self.0.clone()) + } + } + + /// Rewrites the input to whatever it last saw, prefixed — so a second + /// modification proves it observed the first. + struct Appending(&'static str); + + #[async_trait] + impl PreExecutionHook for Appending { + async fn pre_tool_execution( + &self, + context: &PreExecutionContext, + ) -> Result { + Ok(HookDecision::Modify { + input_json: format!("{}{}", context.input_json, self.0), + reason: None, + }) + } + } + + fn context() -> PreExecutionContext { + PreExecutionContext { + agent_id: "a1".to_string(), + tool_name: "shell".to_string(), + tool_call_id: "tc-1".to_string(), + input_json: "start".to_string(), + working_directory: PathBuf::from("/repo"), + } + } + + #[tokio::test] + async fn no_hooks_allows() { + let hooks = PreExecutionHooks::new(); + assert_eq!(hooks.run(&context()).await.unwrap(), HookDecision::Allow); + } + + #[tokio::test] + async fn a_deny_short_circuits_the_rest() { + let hooks = PreExecutionHooks::new() + .with_hook(Fixed(HookDecision::Deny("no".to_string()))) + .with_hook(Appending("-never")); + + assert_eq!( + hooks.run(&context()).await.unwrap(), + HookDecision::Deny("no".to_string()), + "a hook after a denial must not get to overwrite the answer" + ); + } + + #[tokio::test] + async fn modifications_compose_in_order() { + let hooks = PreExecutionHooks::new() + .with_hook(Appending("-one")) + .with_hook(Appending("-two")); + + let HookDecision::Modify { input_json, .. } = hooks.run(&context()).await.unwrap() else { + panic!("expected a modification"); + }; + assert_eq!( + input_json, "start-one-two", + "each hook must see the input as its predecessors left it" + ); + } + + #[tokio::test] + async fn a_later_hook_can_still_deny_a_modified_call() { + let hooks = PreExecutionHooks::new() + .with_hook(Appending("-one")) + .with_hook(Fixed(HookDecision::Deny("still no".to_string()))); + + assert_eq!( + hooks.run(&context()).await.unwrap(), + HookDecision::Deny("still no".to_string()), + "modify must not be a way around a hook that runs later" + ); + } + + /// Awaits before answering, which is the whole reason the trait is async: + /// under the old sync signature this had to be `block_in_place`, and that + /// panics on a current_thread runtime. + struct Awaits; + + #[async_trait] + impl PreExecutionHook for Awaits { + async fn pre_tool_execution( + &self, + _context: &PreExecutionContext, + ) -> Result { + tokio::task::yield_now().await; + tokio::time::sleep(std::time::Duration::from_millis(1)).await; + Ok(HookDecision::Deny("after awaiting".to_string())) + } + } + + #[tokio::test(flavor = "current_thread")] + async fn a_hook_may_await_even_on_a_current_thread_runtime() { + let hooks = PreExecutionHooks::new().with_hook(Awaits); + + assert_eq!( + hooks.run(&context()).await.unwrap(), + HookDecision::Deny("after awaiting".to_string()), + "a hook doing real work must not need a multi-thread runtime" + ); + } +} diff --git a/vendor/mentra/src/runtime/control/policy.rs b/vendor/mentra/src/runtime/control/policy.rs new file mode 100644 index 0000000..559c3e7 --- /dev/null +++ b/vendor/mentra/src/runtime/control/policy.rs @@ -0,0 +1,789 @@ +use std::{ + fs, + path::{Component, Path, PathBuf}, + time::Duration, +}; + +use crate::tool::{ + ToolAuthorizationOutcome, + bash_validation::{CommandIntent, ValidationResult, classify_command, validate_command}, +}; + +/// Controls heuristic validation of builtin shell commands. +/// +/// Shell validation is a defense-in-depth guardrail and permission-prompt UX +/// signal. It is heuristic and is not a security boundary. Working-directory +/// checks do not confine a shell process; filesystem and network isolation +/// require an OS-enforced [`crate::runtime::RuntimeExecutor`]. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub enum ShellValidationMode { + /// Classify commands for authorization previews without changing execution. + #[default] + Off, + /// Emit an authorization hook for warnings or blocks, but allow execution. + Warn, + /// Deny commands classified as blocked and surface warnings through hooks. + Enforce, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ShellValidation { + pub(crate) mode: ShellValidationMode, + pub(crate) intent: CommandIntent, + pub(crate) result: ValidationResult, + pub(crate) outcome: ToolAuthorizationOutcome, +} + +impl ShellValidationMode { + pub(crate) const fn as_str(self) -> &'static str { + match self { + Self::Off => "off", + Self::Warn => "warn", + Self::Enforce => "enforce", + } + } +} + +impl ShellValidation { + pub(crate) const fn intent_name(&self) -> &'static str { + match self.intent { + CommandIntent::ReadOnly => "read_only", + CommandIntent::Write => "write", + CommandIntent::Destructive => "destructive", + CommandIntent::Network => "network", + CommandIntent::ProcessManagement => "process_management", + CommandIntent::PackageManagement => "package_management", + CommandIntent::SystemAdmin => "system_admin", + CommandIntent::Unknown => "unknown", + } + } + + pub(crate) fn reason(&self) -> Option<&str> { + self.result.reason() + } + + pub(crate) fn should_emit_hook(&self) -> bool { + self.mode != ShellValidationMode::Off && self.outcome != ToolAuthorizationOutcome::Allow + } + + pub(crate) fn should_deny(&self) -> bool { + self.mode == ShellValidationMode::Enforce && self.outcome == ToolAuthorizationOutcome::Deny + } +} + +/// Authorization policy for builtin shell, background, and file tools. +#[derive(Debug, Clone)] +pub struct RuntimePolicy { + allow_shell_commands: bool, + allow_background_commands: bool, + allowed_working_roots: Vec, + allowed_read_roots: Vec, + allowed_write_roots: Vec, + denied_write_roots: Vec, + allowed_env_vars: Vec, + shell_validation_mode: ShellValidationMode, + pub(crate) background_task_limit: Option, + pub(crate) default_command_timeout: Duration, + pub(crate) max_command_timeout: Duration, + pub(crate) max_output_bytes_per_stream: usize, + pub(crate) max_tool_result_bytes: usize, + pub(crate) max_tool_result_lines: usize, + pub(crate) spill_full_tool_output: bool, +} + +impl Default for RuntimePolicy { + fn default() -> Self { + Self { + allow_shell_commands: false, + allow_background_commands: false, + allowed_working_roots: Vec::new(), + allowed_read_roots: Vec::new(), + allowed_write_roots: Vec::new(), + denied_write_roots: Vec::new(), + allowed_env_vars: default_allowed_env_vars(), + shell_validation_mode: ShellValidationMode::Off, + background_task_limit: Some(8), + default_command_timeout: Duration::from_secs(30), + max_command_timeout: Duration::from_secs(30), + max_output_bytes_per_stream: 64 * 1024, + max_tool_result_bytes: 50 * 1024, + max_tool_result_lines: 2_000, + spill_full_tool_output: true, + } + } +} + +fn default_allowed_env_vars() -> Vec { + #[cfg(windows)] + { + let mut vars = vec!["PATH".to_string()]; + vars.extend([ + "PATHEXT".to_string(), + "SystemRoot".to_string(), + "COMSPEC".to_string(), + "TEMP".to_string(), + "TMP".to_string(), + ]); + vars + } + + #[cfg(not(windows))] + { + vec!["PATH".to_string()] + } +} + +impl RuntimePolicy { + /// Returns a permissive policy that enables shell and background execution. + pub fn permissive() -> Self { + Self { + allow_shell_commands: true, + allow_background_commands: true, + ..Self::default() + } + } + + /// Returns a workspace-bounded policy for builtin file tools. + /// + /// Shell and background execution remain disabled because the builtin + /// local executor runs directly on the host and a working-directory check + /// cannot confine filesystem or network effects. Hosts that install an + /// OS-enforced executor through [`crate::runtime::Runtime::builder`] may + /// explicitly opt in with [`Self::allow_shell_commands`] and + /// [`Self::allow_background_commands`]. + pub fn workspace_bounded(workspace: impl Into) -> Self { + let workspace = workspace.into(); + Self { + allowed_working_roots: vec![workspace.clone()], + allowed_read_roots: vec![workspace.clone()], + allowed_write_roots: vec![workspace], + default_command_timeout: Duration::from_secs(120), + max_command_timeout: Duration::from_secs(600), + ..Self::default() + } + } + + /// Returns a policy that allows builtin file reads but blocks builtin file + /// writes and host shell execution. + /// + /// A host may opt into shell execution only after installing an executor + /// that enforces read-only filesystem and network policy at the OS boundary. + pub fn read_only(workspace: impl Into) -> Self { + let workspace = workspace.into(); + Self { + allow_background_commands: false, + allowed_working_roots: vec![workspace.clone()], + allowed_read_roots: vec![workspace], + allowed_write_roots: Vec::new(), + ..Self::default() + } + } + + /// Enables or disables foreground shell command execution. + /// + /// This switch grants authority to the configured executor; it does not + /// sandbox the builtin `LocalRuntimeExecutor`. + pub fn allow_shell_commands(mut self, allow: bool) -> Self { + self.allow_shell_commands = allow; + self + } + + /// Enables or disables background shell command execution. + /// + /// This switch grants authority to the configured executor; it does not + /// sandbox the builtin `LocalRuntimeExecutor`. + pub fn allow_background_commands(mut self, allow: bool) -> Self { + self.allow_background_commands = allow; + self + } + + /// Selects heuristic validation for builtin shell commands. + /// + /// This is a defense-in-depth guardrail and prompt signal, not a security + /// boundary. [`ShellValidationMode::Off`] preserves execution behavior. + pub fn shell_validation(mut self, mode: ShellValidationMode) -> Self { + self.shell_validation_mode = mode; + self + } + + /// Adds an extra working-directory root allowed for shell commands. + pub fn with_allowed_working_root(mut self, path: impl Into) -> Self { + self.allowed_working_roots.push(path.into()); + self + } + + /// Adds an extra root allowed for builtin file reads. + pub fn with_allowed_read_root(mut self, path: impl Into) -> Self { + self.allowed_read_roots.push(path.into()); + self + } + + /// Adds an extra root allowed for builtin file writes. + pub fn with_allowed_write_root(mut self, path: impl Into) -> Self { + self.allowed_write_roots.push(path.into()); + self + } + + /// Carves a hole in the write roots: a path under `path` is refused even + /// when an allow-root would otherwise permit it. + /// + /// For the places inside a workspace that an agent should not be able to + /// change because changing them changes what runs — `.git/hooks` being the + /// canonical one, since a file written there executes on the next commit. + /// Allow-roots alone cannot express it: the whole workspace is writable and + /// these are inside the workspace. + /// + /// **This binds mentra's builtin file tools, not the shell.** A command + /// like `sh -c 'echo … > .git/hooks/pre-commit'` still reaches the path, + /// because the runtime does not parse shell and cannot know where a + /// redirect points. Treat this as hygiene that closes the obvious route, + /// never as a boundary — the boundary belongs to the OS. + pub fn with_denied_write_root(mut self, path: impl Into) -> Self { + self.denied_write_roots.push(path.into()); + self + } + + /// Records an environment variable name that callers may expose to tools. + pub fn with_allowed_env_var(mut self, name: impl Into) -> Self { + self.allowed_env_vars.push(name.into()); + self + } + + /// Sets the maximum number of concurrently tracked background tasks per agent. + pub fn with_max_background_tasks(mut self, limit: usize) -> Self { + self.background_task_limit = Some(limit); + self + } + + /// Sets the default builtin command timeout. + pub fn with_default_command_timeout(mut self, timeout: Duration) -> Self { + self.default_command_timeout = timeout; + self + } + + /// Sets the hard timeout cap for builtin commands. + pub fn with_max_command_timeout(mut self, timeout: Duration) -> Self { + self.max_command_timeout = timeout; + self + } + + /// Sets the maximum captured bytes for each output stream. + pub fn with_max_output_bytes_per_stream(mut self, max_bytes: usize) -> Self { + self.max_output_bytes_per_stream = max_bytes; + self + } + + /// Sets the provider-visible byte limit for each completed tool result. + /// + /// The limit applies independently to successful and error results. An + /// actionable truncation notice is appended outside the retained head. + pub fn with_max_tool_result_bytes(mut self, max_bytes: usize) -> Self { + self.max_tool_result_bytes = max_bytes; + self + } + + /// Sets the provider-visible line limit for each completed tool result. + pub fn with_max_tool_result_lines(mut self, max_lines: usize) -> Self { + self.max_tool_result_lines = max_lines; + self + } + + /// Enables or disables spilling a truncated tool result to the agent's + /// transcript artifact directory. + pub fn spill_full_tool_output(mut self, spill: bool) -> Self { + self.spill_full_tool_output = spill; + self + } + + /// Backward-compatible shortcut that sets both default and max timeout. + pub fn with_command_timeout(mut self, timeout: Duration) -> Self { + self.default_command_timeout = timeout; + self.max_command_timeout = timeout; + self + } + + pub(crate) fn authorize_command_execution( + &self, + base_dir: &Path, + cwd: &Path, + background: bool, + ) -> Result<(), String> { + self.authorize_command_roots(base_dir, cwd, background) + } + + pub(crate) fn evaluate_shell_command( + &self, + command: &str, + default_workspace: &Path, + ) -> ShellValidation { + let workspace = self + .allowed_working_roots + .first() + .map(PathBuf::as_path) + .unwrap_or(default_workspace); + let result = validate_command(command, workspace, self.allowed_write_roots.is_empty()); + let outcome = result.authorization_outcome(); + + ShellValidation { + mode: self.shell_validation_mode, + intent: classify_command(command), + result, + outcome, + } + } + + pub(crate) fn effective_timeout(&self, requested: Option) -> Duration { + requested + .unwrap_or(self.default_command_timeout) + .min(self.max_command_timeout) + } + + pub(crate) fn allowed_environment(&self) -> Vec<(String, String)> { + self.allowed_env_vars + .iter() + .filter_map(|name| std::env::var(name).ok().map(|value| (name.clone(), value))) + .collect() + } + + pub(crate) fn authorize_file_read( + &self, + base_dir: &Path, + path: &Path, + ) -> Result { + let resolved = resolve_authorized_path(base_dir, path)?; + + if path_is_allowed( + resolved.as_path(), + base_dir, + self.allowed_read_roots.as_slice(), + ) { + Ok(resolved) + } else { + Err(format!( + "Path '{}' is outside the runtime policy read roots", + resolved.display() + )) + } + } + + pub(crate) fn authorize_file_write( + &self, + base_dir: &Path, + path: &Path, + ) -> Result { + let resolved = resolve_authorized_path(base_dir, path)?; + + // Checked before the allow list, because a denial is only meaningful + // inside a root that would otherwise permit the write. Both sides + // normalize through `normalize_policy_root`, so `.git/hooks/../hooks` + // and a symlink into a denied root resolve to the same answer. + if path_is_under_any(resolved.as_path(), self.denied_write_roots.as_slice()) { + return Err(format!( + "Path '{}' is inside a runtime policy denied write root", + resolved.display() + )); + } + + if path_is_allowed( + resolved.as_path(), + base_dir, + self.allowed_write_roots.as_slice(), + ) { + Ok(resolved) + } else { + Err(format!( + "Path '{}' is outside the runtime policy write roots", + resolved.display() + )) + } + } + + fn authorize_command_roots( + &self, + base_dir: &Path, + cwd: &Path, + background: bool, + ) -> Result<(), String> { + if !self.allow_shell_commands { + return Err( + "Shell command execution is disabled by the runtime policy. Use RuntimeBuilder::with_policy(...) to opt in." + .to_string(), + ); + } + if background && !self.allow_background_commands { + return Err( + "Background command execution is disabled by the runtime policy.".to_string(), + ); + } + + if !path_is_allowed(cwd, base_dir, self.allowed_working_roots.as_slice()) { + return Err(format!( + "Working directory '{}' is outside the runtime policy roots", + cwd.display() + )); + } + + Ok(()) + } +} + +/// Whether `path` sits under any of `roots`, comparing resolved forms. +fn path_is_under_any(path: &Path, roots: &[PathBuf]) -> bool { + if roots.is_empty() { + return false; + } + let candidate = normalize_policy_root(path); + roots + .iter() + .map(|root| normalize_policy_root(root)) + .any(|root| candidate.starts_with(root)) +} + +fn path_is_allowed(path: &Path, default_root: &Path, extra_roots: &[PathBuf]) -> bool { + let candidate_path = normalize_policy_root(path); + let default_root = normalize_policy_root(default_root); + candidate_path.starts_with(&default_root) + || extra_roots + .iter() + .map(|root| normalize_policy_root(root)) + .any(|root| candidate_path.starts_with(root)) +} + +fn normalize_policy_root(path: &Path) -> PathBuf { + normalize_absolute_path(path) + .ok() + .and_then(|normalized| resolve_existing_components(&normalized).ok()) + .unwrap_or_else(|| fs::canonicalize(path).unwrap_or_else(|_| path.to_path_buf())) +} + +fn resolve_authorized_path(base_dir: &Path, path: &Path) -> Result { + let resolved = if path.is_absolute() { + path.to_path_buf() + } else { + base_dir.join(path) + }; + let normalized = normalize_absolute_path(&resolved)?; + resolve_existing_components(&normalized) +} + +fn normalize_absolute_path(path: &Path) -> Result { + if !path.is_absolute() { + return Err(format!( + "Path '{}' must resolve to an absolute path", + path.display() + )); + } + + let mut normalized = PathBuf::new(); + for component in path.components() { + match component { + Component::Prefix(prefix) => normalized.push(prefix.as_os_str()), + Component::RootDir => normalized.push(component.as_os_str()), + Component::CurDir => {} + Component::ParentDir => { + if !normalized.pop() || !normalized.is_absolute() { + return Err(format!( + "Path '{}' escapes the filesystem root", + path.display() + )); + } + } + Component::Normal(segment) => normalized.push(segment), + } + } + + if !normalized.is_absolute() { + return Err(format!( + "Path '{}' must resolve to an absolute path", + path.display() + )); + } + + Ok(normalized) +} + +fn resolve_existing_components(path: &Path) -> Result { + let mut resolved = PathBuf::new(); + for component in path.components() { + match component { + Component::Prefix(prefix) => resolved.push(prefix.as_os_str()), + Component::RootDir => resolved.push(component.as_os_str()), + Component::CurDir => {} + Component::ParentDir => unreachable!("paths are normalized before resolution"), + Component::Normal(segment) => { + resolved.push(segment); + match fs::symlink_metadata(&resolved) { + Ok(metadata) if metadata.file_type().is_symlink() => { + resolved = fs::canonicalize(&resolved).map_err(|error| { + format!( + "Failed to resolve symlink '{}': {error}", + resolved.display() + ) + })?; + } + Ok(_) => {} + Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} + Err(error) => { + return Err(format!( + "Failed to inspect '{}': {error}", + resolved.display() + )); + } + } + } + } + } + + if !resolved.is_absolute() { + return Err(format!( + "Path '{}' must resolve to an absolute path", + path.display() + )); + } + + Ok(resolved) +} + +#[cfg(test)] +mod tests { + use super::*; + #[cfg(unix)] + use std::{ + fs, + time::{SystemTime, UNIX_EPOCH}, + }; + + fn test_path(label: &str) -> PathBuf { + std::env::temp_dir() + .join("mentra-runtime-policy-tests") + .join(label) + } + + #[test] + fn shell_roots_and_background_switches_short_circuit() { + let cwd = test_path("repo"); + let policy = RuntimePolicy::default() + .allow_shell_commands(true) + .allow_background_commands(false); + let error = policy + .authorize_command_execution(&cwd, &cwd, true) + .expect_err("background should be disabled"); + assert!(error.contains("Background command execution is disabled")); + } + + #[test] + fn bounded_policies_keep_host_shell_execution_disabled() { + let workspace = test_path("bounded-repo"); + + for policy in [ + RuntimePolicy::workspace_bounded(&workspace), + RuntimePolicy::read_only(&workspace), + ] { + let error = policy + .authorize_command_execution(&workspace, &workspace, false) + .expect_err("bounded policy must not authorize the host shell"); + assert!(error.contains("Shell command execution is disabled")); + } + } + + #[test] + fn bounded_policy_allows_explicit_external_executor_opt_in() { + let workspace = test_path("sandboxed-repo"); + let policy = RuntimePolicy::workspace_bounded(&workspace) + .allow_shell_commands(true) + .allow_background_commands(true); + + policy + .authorize_command_execution(&workspace, &workspace, false) + .expect("foreground shell opt-in"); + policy + .authorize_command_execution(&workspace, &workspace, true) + .expect("background shell opt-in"); + } + + #[test] + fn shell_validation_defaults_off_and_uses_authorization_semantics() { + let workspace = test_path("validation-workspace"); + let default_validation = + RuntimePolicy::default().evaluate_shell_command("rm -rf /tmp/sentinel", &workspace); + assert_eq!(default_validation.mode, ShellValidationMode::Off); + assert_eq!(default_validation.intent, CommandIntent::Destructive); + assert_eq!(default_validation.outcome, ToolAuthorizationOutcome::Deny); + assert!(!default_validation.should_deny()); + + let warned = RuntimePolicy::default() + .shell_validation(ShellValidationMode::Warn) + .evaluate_shell_command("rm -rf /tmp/sentinel", &workspace); + assert!(warned.should_emit_hook()); + assert!(!warned.should_deny()); + + let enforced = RuntimePolicy::default() + .shell_validation(ShellValidationMode::Enforce) + .evaluate_shell_command("rm -rf /tmp/sentinel", &workspace); + assert!(enforced.should_emit_hook()); + assert!(enforced.should_deny()); + + let enforced_warning = RuntimePolicy::workspace_bounded(&workspace) + .shell_validation(ShellValidationMode::Enforce) + .evaluate_shell_command("rm -rf /", &workspace); + assert_eq!(enforced_warning.outcome, ToolAuthorizationOutcome::Prompt); + assert!(enforced_warning.should_emit_hook()); + assert!(!enforced_warning.should_deny()); + } + + #[test] + fn tool_result_limits_have_stable_defaults_and_builders() { + let defaults = RuntimePolicy::default(); + assert_eq!(defaults.max_tool_result_bytes, 50 * 1024); + assert_eq!(defaults.max_tool_result_lines, 2_000); + assert!(defaults.spill_full_tool_output); + + let configured = defaults + .with_max_tool_result_bytes(123) + .with_max_tool_result_lines(7) + .spill_full_tool_output(false); + assert_eq!(configured.max_tool_result_bytes, 123); + assert_eq!(configured.max_tool_result_lines, 7); + assert!(!configured.spill_full_tool_output); + } + + #[test] + fn authorize_command_execution_rejects_working_directory_outside_roots() { + let base_dir = test_path("repo"); + let cwd = test_path("other"); + let policy = RuntimePolicy::default().allow_shell_commands(true); + + let error = policy + .authorize_command_execution(&base_dir, &cwd, false) + .expect_err("working directory should be rejected"); + assert!(error.contains("outside the runtime policy roots")); + } + + #[test] + fn normalize_absolute_path_rejects_parent_past_root() { + let mut path = std::env::temp_dir(); + for _ in 0..10 { + path.push(".."); + } + path.push("escape"); + let error = normalize_absolute_path(&path).expect_err("path should be rejected"); + assert!(error.contains("escapes the filesystem root")); + } + + #[cfg(unix)] + #[test] + fn authorize_file_write_rejects_symlink_escape() { + use std::os::unix::fs::symlink; + + let root = unique_temp_dir("policy-write-root"); + let outside = unique_temp_dir("policy-write-outside"); + let link = root.join("link"); + symlink(&outside, &link).expect("create symlink"); + + let policy = RuntimePolicy::default().with_allowed_write_root(&root); + let error = policy + .authorize_file_write(&root, &link.join("escape.txt")) + .expect_err("symlink escape should be denied"); + assert!(error.contains("outside the runtime policy write roots")); + + let _ = fs::remove_dir_all(&root); + let _ = fs::remove_dir_all(&outside); + } + + #[cfg(unix)] + #[test] + fn a_denied_root_inside_an_allowed_one_refuses_the_write() { + let root = unique_temp_dir("policy-deny-root"); + let hooks = root.join(".git").join("hooks"); + fs::create_dir_all(&hooks).expect("create hooks dir"); + + let policy = RuntimePolicy::default() + .with_allowed_write_root(&root) + .with_denied_write_root(&hooks); + + // The whole point: the workspace is writable and this is inside it, so + // allow-roots alone could never express the carve-out. + let error = policy + .authorize_file_write(&root, &hooks.join("pre-commit")) + .expect_err("a denied root must win over the allow root containing it"); + assert!(error.contains("denied write root"), "got: {error}"); + + // A sibling under the same allow root is untouched. + policy + .authorize_file_write(&root, &root.join("src.rs")) + .expect("an ordinary write is unaffected"); + + let _ = fs::remove_dir_all(&root); + } + + #[cfg(unix)] + #[test] + fn a_traversal_into_a_denied_root_is_refused() { + let root = unique_temp_dir("policy-deny-traverse"); + let hooks = root.join(".git").join("hooks"); + fs::create_dir_all(&hooks).expect("create hooks dir"); + + let policy = RuntimePolicy::default() + .with_allowed_write_root(&root) + .with_denied_write_root(&hooks); + + // Spelled to look like it lands elsewhere. Both sides normalize, so + // the spelling does not decide the answer. + let sneaky = root.join(".git").join("hooks").join("..").join("hooks"); + let error = policy + .authorize_file_write(&root, &sneaky.join("pre-push")) + .expect_err("a path that resolves into a denied root is still denied"); + assert!(error.contains("denied write root"), "got: {error}"); + + let _ = fs::remove_dir_all(&root); + } + + #[cfg(unix)] + #[test] + fn a_symlink_into_a_denied_root_is_refused() { + use std::os::unix::fs::symlink; + + let root = unique_temp_dir("policy-deny-symlink"); + let hooks = root.join(".git").join("hooks"); + fs::create_dir_all(&hooks).expect("create hooks dir"); + let link = root.join("shortcut"); + symlink(&hooks, &link).expect("create symlink"); + + let policy = RuntimePolicy::default() + .with_allowed_write_root(&root) + .with_denied_write_root(&hooks); + + let error = policy + .authorize_file_write(&root, &link.join("pre-commit")) + .expect_err("a symlink is not a way around a denied root"); + assert!(error.contains("denied write root"), "got: {error}"); + + let _ = fs::remove_dir_all(&root); + } + + #[cfg(unix)] + #[test] + fn no_denied_roots_changes_nothing() { + let root = unique_temp_dir("policy-deny-empty"); + fs::create_dir_all(&root).expect("create root"); + + let policy = RuntimePolicy::default().with_allowed_write_root(&root); + + policy + .authorize_file_write(&root, &root.join(".git").join("hooks").join("pre-commit")) + .expect("with no deny list the allow root decides alone"); + + let _ = fs::remove_dir_all(&root); + } + + #[cfg(unix)] + fn unique_temp_dir(label: &str) -> PathBuf { + let unique = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("duration") + .as_nanos(); + let path = std::env::temp_dir().join(format!("mentra-{label}-{unique}")); + fs::create_dir_all(&path).expect("create temp dir"); + path + } +} diff --git a/vendor/mentra/src/runtime/control/run.rs b/vendor/mentra/src/runtime/control/run.rs new file mode 100644 index 0000000..85d4d9e --- /dev/null +++ b/vendor/mentra/src/runtime/control/run.rs @@ -0,0 +1,666 @@ +use std::sync::{ + Arc, OnceLock, + atomic::{AtomicBool, AtomicU64, Ordering}, +}; +use std::time::{Duration, SystemTime}; + +use crate::runtime::error::RuntimeError; + +const DEFAULT_PROVIDER_RETRY_BUDGET: usize = 5; +const DEFAULT_PROVIDER_RETRY_BASE_DELAY: Duration = Duration::from_millis(500); +const DEFAULT_PROVIDER_RETRY_MAX_DELAY: Duration = Duration::from_secs(5); +const DEFAULT_PROVIDER_RETRY_AFTER_CAP: Duration = Duration::from_secs(60); + +/// How long a run waits between provider retries. +/// +/// *How many* retries it gets stays on [`RunOptions::retry_budget`]; see that +/// field for why the count did not move in here. This is the other half: what +/// each of those attempts waits for. +/// +/// The defaults reproduce mentra's historical schedule exactly — 500 ms, +/// doubling, capped at 5 s — which is shaped for a blip: a connection reset, a +/// tunnel restart, a 502 from a proxy that is already coming back. A rate limit +/// is a different failure. It lasts as long as the window it belongs to, which +/// is routinely a minute, and the whole default budget elapses in about twelve +/// and a half seconds — five attempts into a limit that was never going to lift +/// in that time, and then a lost turn. A host that knows it is behind a metered +/// gateway can say so here instead of living with a schedule chosen for a +/// different failure. +/// +/// ```rust +/// use std::time::Duration; +/// use mentra::runtime::{ProviderRetry, RunOptions}; +/// +/// // Wait out a minute-long rate-limit window rather than a blip. +/// let options = RunOptions { +/// retry_budget: 8, +/// ..RunOptions::default() +/// } +/// .with_provider_retry(ProviderRetry { +/// base_delay: Duration::from_secs(1), +/// max_delay: Duration::from_secs(30), +/// ..ProviderRetry::default() +/// }); +/// # let _ = options; +/// ``` +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ProviderRetry { + /// The wait before the second attempt, doubled before each attempt after + /// it. + pub base_delay: Duration, + /// The ceiling the doubling stops at. Reached and then held, so a long + /// budget spends its tail attempts at a steady interval rather than an + /// ever-growing one. + pub max_delay: Duration, + /// The longest wait a *server* may impose through `Retry-After`. + /// + /// A server that answers `Retry-After: 3600` is not describing a rate + /// limit any run should sit through, and honoring it unconditionally hands + /// a remote party control of how long this process blocks. The header is + /// clamped to this before it is considered. It never shortens + /// [`max_delay`](Self::max_delay): a schedule the host chose is the host's + /// business, and this bounds only what the other end asked for. + pub retry_after_cap: Duration, +} + +impl Default for ProviderRetry { + fn default() -> Self { + Self { + base_delay: DEFAULT_PROVIDER_RETRY_BASE_DELAY, + max_delay: DEFAULT_PROVIDER_RETRY_MAX_DELAY, + retry_after_cap: DEFAULT_PROVIDER_RETRY_AFTER_CAP, + } + } +} + +impl ProviderRetry { + /// The wait this schedule prescribes before the attempt that follows + /// `attempt` (one-based), before anything the provider said is considered. + pub fn scheduled_delay(&self, attempt: usize) -> Duration { + // Clamped to `u32`'s width, not `usize`'s: the factor is a `u32`, and + // shifting one by 32 or more is a panic in debug and nonsense in + // release. Unreachable at the default budget of five, reachable the + // moment a host raises it — which is the point of this type. + let shift = attempt.saturating_sub(1).min(u32::BITS as usize - 1) as u32; + let factor = 1u32 << shift; + self.base_delay + .checked_mul(factor) + .unwrap_or(self.max_delay) + .min(self.max_delay) + } + + /// The wait actually taken before the attempt that follows `attempt`, given + /// what the provider asked for in `retry_after`. + /// + /// The longer of the two wins, because they answer different questions: the + /// schedule is the host's floor on how hard it is willing to hammer a + /// provider, and `Retry-After` is the provider's floor on when it will + /// answer again. Waiting the shorter of them satisfies neither. The + /// server's number is clamped to + /// [`retry_after_cap`](Self::retry_after_cap) first. + pub fn delay_for(&self, attempt: usize, retry_after: Option) -> Duration { + let scheduled = self.scheduled_delay(attempt); + match retry_after { + Some(requested) => scheduled.max(requested.min(self.retry_after_cap)), + None => scheduled, + } + } +} + +/// Why a run ended before its work was done, when a bound rather than the model +/// decided that. +/// +/// The two graceful bounds — [`RunOptions::stop`] and +/// [`RunOptions::token_budget`] — deliberately end a run the way the model +/// finishing does: at a round boundary, transcript committed, `Ok`. That is the +/// right *behavior* and a silent *report*, because what the caller receives is +/// then identical for "the model was done" and "the runner refused to start +/// another round". A caller that has to tell those apart — a CLI owing a +/// distinct exit code, a supervisor deciding whether to prompt again — is +/// otherwise left either recomputing the comparison the runner already made or +/// reading prose, and the first answers a slightly different question (what is +/// true *now*, not what was true at the boundary) while the second is not a +/// contract at all. This records the runner's own decision at the moment it +/// made it. +/// +/// Read it back through [`RunOptions::ended_early`]. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub enum EarlyEnd { + /// The run ended at a round boundary because [`RunOptions::stop`] was + /// tripped. + /// + /// Reported in preference to [`TokenBudget`](Self::TokenBudget) when both + /// were true at that boundary. A stop is an instruction the caller issued, + /// and the runner would have ended there with no budget set at all; a + /// crossed budget is an ambient bound that merely also held. Naming the + /// bound would tell a caller its allowance ran out when what actually + /// happened is that it asked to stop — and the runner's own control flow + /// agrees, since it checks `stop` first and never reaches the budget check. + StopRequested, + /// The run ended at a round boundary because cumulative reported usage had + /// reached or passed [`RunOptions::token_budget`]. + /// + /// See that field for why this can only ever be noticed at a boundary, and + /// for what the run keeps when it is. + TokenBudget, +} + +/// A shared flag a caller trips to stop a run. +/// +/// `Debug` prints whether it has been tripped, so a host that embeds one in +/// its own options struct can still derive `Debug` on that struct. +#[derive(Clone, Default)] +pub struct CancellationToken { + cancelled: Arc, +} + +impl std::fmt::Debug for CancellationToken { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("CancellationToken") + .field("cancelled", &self.is_cancelled()) + .finish() + } +} + +pub type CancellationFlag = CancellationToken; + +impl CancellationToken { + pub fn cancel(&self) { + self.cancelled.store(true, Ordering::SeqCst); + } + + pub fn is_cancelled(&self) -> bool { + self.cancelled.load(Ordering::SeqCst) + } +} + +#[derive(Clone)] +pub struct RunOptions { + pub cancellation: Option, + /// A **graceful** stop signal, distinct from [`cancellation`](Self::cancellation). + /// + /// When this token is tripped (via [`CancellationToken::cancel`]) the run ends + /// **successfully** at the next round boundary — the committed transcript is + /// kept (the run resolves like the model self-terminating with no further tool + /// calls), rather than failing and rolling the run back the way `cancellation` + /// does. Use it to stop gathering once enough work is done while preserving the + /// gathered context for a follow-up turn on the same agent. `None` (the default) + /// never stops the run. + pub stop: Option, + pub deadline: Option, + /// How many times one provider request may be re-attempted after a + /// transient failure, before the run gives up and reports the error. + /// + /// The count stayed here rather than moving into [`ProviderRetry`] beside + /// the schedule it belongs with, because + /// `RunOptions { retry_budget: 3, ..default() }` is how every host that has + /// ever changed this wrote it, and there is no spelling of a moved public + /// field that keeps those compiling. One number, one home; + /// [`provider_retry`](Self::provider_retry) holds the rest. + /// + /// **Retries are model requests.** Each attempt increments the same counter + /// [`model_budget`](Self::model_budget) bounds, so a run with both set can + /// exhaust its model budget on retries and end in + /// [`ModelBudgetExceeded`](crate::error::RuntimeError::ModelBudgetExceeded) + /// without the model ever having answered. That is deliberate: + /// `model_budget` bounds how many times this run may reach for the + /// provider, and an attempt that failed still reached. A host raising this + /// budget to sit out a rate limit should raise `model_budget` with it, or + /// leave `model_budget` at `None`, where no such interaction exists. + /// + /// Inherited by [`child`](Self::child) runs, along with + /// [`provider_retry`](Self::provider_retry); see that method for why a + /// delegated run meets the same provider with the same patience. + pub retry_budget: usize, + /// How long each of those retries waits. See [`ProviderRetry`]; the default + /// schedule is mentra's historical one, unchanged. + pub provider_retry: ProviderRetry, + pub tool_budget: Option, + /// A bound on how many provider requests this run may make, counting failed + /// attempts — see [`retry_budget`](Self::retry_budget). `None` (the + /// default) never bounds the run. + pub model_budget: Option, + /// A per-run [`RoundStrategy`](crate::agent::RoundStrategy) invoked at each + /// round boundary (after a committed tool round and after a committed + /// tool-free assistant message). It is owned by this single `Agent::run` + /// invocation, never by a shared [`Runtime`](crate::Runtime), so one run's + /// steering and stop state cannot leak into another run. `None` (the default) + /// reproduces mentra's built-in round loop exactly. + pub round_strategy: Option>, + /// A **soft** aggregate token bound on this run's reported usage, distinct + /// from [`model_budget`](Self::model_budget) (which caps the number of + /// provider *requests*, not tokens). + /// + /// Token usage is only known once a round's response has streamed in full + /// (the same point where `TurnRunner` emits + /// `AgentEvent::UsageReport`), so this can never be a hard ceiling: a single + /// round is always allowed to finish even if it pushes cumulative usage from + /// under the bound to well past it. Once a round has completed, the bound is + /// checked at the same round-boundary point where [`stop`](Self::stop) is + /// checked: if cumulative reported `input_tokens + output_tokens` (summed + /// across every round this run, and any [`child`](Self::child) run sharing + /// this handle, has completed) has reached or exceeded the bound, the run + /// ends **gracefully** there, exactly as `stop` does — the committed + /// transcript is kept, not rolled back. Cache-read and cache-creation tokens + /// are not counted. `None` (the default) never stops the run. This is never + /// an expense bound: mentra has no injected price source and makes no + /// monetary claim. + pub token_budget: Option, + /// Shared cumulative `input_tokens + output_tokens` counter backing + /// [`token_budget`](Self::token_budget) and read back through + /// [`reported_tokens`](Self::reported_tokens). Held behind an `Arc` so a + /// [`child`](Self::child) run reports into the same aggregate as its parent — + /// that is the intended way to share it. This field is `pub` only so + /// `RunOptions { .., ..RunOptions::default() }` construction keeps working; + /// leave it at its default (a fresh, zeroed counter) unless you are + /// deliberately aliasing a specific run's accounting. + pub token_usage: Arc, + /// Where a run records *why* it ended early, read back through + /// [`ended_early`](Self::ended_early). + /// + /// The counterpart of [`token_usage`](Self::token_usage) for the decision + /// rather than the count: the runner knows at the boundary it stops at + /// which bound stopped it, and this is how that reaches a caller holding a + /// clone of these options instead of being re-derived — or lost. Written at + /// most once, first writer winning, because a run ends at exactly one + /// boundary and because both conditions that produce an entry are sticky + /// under a fixed bound: a tripped stop token stays tripped, and a crossed + /// cumulative total stays crossed, so a later turn on this handle ends the + /// same way and would record the same answer. Raising + /// [`token_budget`](Self::token_budget) on a clone is the one way to make + /// the entry stale — the deliberate aliasing `token_usage` warns about, + /// seen from the reporting side. Like that field, this one is `pub` only so + /// `RunOptions { .., ..RunOptions::default() }` construction keeps working; + /// leave it at its default (a fresh, empty slot). + pub early_end: Arc>, +} + +impl Default for RunOptions { + fn default() -> Self { + Self { + cancellation: None, + stop: None, + deadline: None, + retry_budget: DEFAULT_PROVIDER_RETRY_BUDGET, + provider_retry: ProviderRetry::default(), + tool_budget: None, + model_budget: None, + round_strategy: None, + token_budget: None, + token_usage: Arc::new(AtomicU64::new(0)), + early_end: Arc::new(OnceLock::new()), + } + } +} + +impl RunOptions { + /// Attaches a per-run [`RoundStrategy`](crate::agent::RoundStrategy) to these + /// options, returning the updated value. + pub fn with_round_strategy(mut self, strategy: Arc) -> Self { + self.round_strategy = Some(strategy); + self + } + + /// Sets the provider retry schedule on these options, returning the updated + /// value. Leaves [`retry_budget`](Self::retry_budget), which counts the + /// attempts this schedule spaces, alone. + pub fn with_provider_retry(mut self, provider_retry: ProviderRetry) -> Self { + self.provider_retry = provider_retry; + self + } + + /// Derives [`RunOptions`] for work spawned during this run — a subagent or a + /// delegated run — sharing this run's aggregate safety bounds: the same + /// [`cancellation`](Self::cancellation) and [`stop`](Self::stop) tokens (so + /// cancelling or gracefully stopping the parent also ends the child), the same + /// [`deadline`](Self::deadline), and the same [`token_budget`](Self::token_budget) + /// bound backed by the *same* accounting handle — a child's reported usage + /// adds to the parent's running total, so parent and child together trip one + /// shared bound rather than each getting an independent one. + /// + /// [`retry_budget`](Self::retry_budget) and + /// [`provider_retry`](Self::provider_retry) carry too, for a different + /// reason: they are not an allowance either run spends but a description of + /// the provider both of them dial. How long that endpoint's rate-limit + /// window lasts does not change because the caller delegated, so a child + /// that reset them would meet the same limiter with the schedule the parent + /// had already found too short — and a host would have to restate its own + /// policy at every delegation boundary to prevent it. Unlike the bounds + /// above they aggregate nothing: a child's retries are its own, and being + /// patient costs the parent none of them. `deadline` and `cancellation`, + /// which do carry, remain the bound on how long all that patience may take. + /// + /// The rest (`tool_budget`, `model_budget`, `round_strategy`) resets to + /// [`RunOptions::default`]: those express per-run policy a child sets + /// independently, bounding work it does rather than describing what it + /// talks to. + /// + /// [`early_end`](Self::early_end) resets too, for a different reason — it + /// records a decision rather than carrying a bound. A child that ends on the + /// shared budget ended *its own* run at *its own* boundary; the parent then + /// reaches its next boundary and records for itself, so keeping the slots + /// apart loses nothing, while sharing one would let a child's ending be read + /// as the parent's — including on a parent that went on to finish its work + /// normally. + /// + /// mentra applies this itself on exactly one path: the `task` intrinsic's + /// delegated subagent runs on the parent run's derived child options, so a + /// model that delegates work cannot spend outside the bounds its own run was + /// given. Every other subagent path is host-driven — call this when + /// threading `RunOptions` into a subagent's or delegated run's own + /// `Agent::run`/`resume` call, including through + /// [`Session::spawn_subagent_with_options`](crate::Session::spawn_subagent_with_options) + /// and, for a custom tool that spawns its own subagent, + /// [`ToolContext::child_run_options`](crate::tool::ToolContext::child_run_options). + pub fn child(&self) -> RunOptions { + RunOptions { + cancellation: self.cancellation.clone(), + stop: self.stop.clone(), + deadline: self.deadline, + token_budget: self.token_budget, + token_usage: Arc::clone(&self.token_usage), + retry_budget: self.retry_budget, + provider_retry: self.provider_retry, + ..RunOptions::default() + } + } + + /// Cumulative `input_tokens + output_tokens` reported so far against + /// [`token_budget`](Self::token_budget), aggregated across this run and any + /// [`child`](Self::child) run sharing this handle. + pub fn reported_tokens(&self) -> u64 { + self.token_usage.load(Ordering::SeqCst) + } + + pub(crate) fn record_tokens(&self, tokens: u64) { + self.token_usage.fetch_add(tokens, Ordering::SeqCst); + } + + /// Why the run ended early, or `None` when none did — the model finished, or + /// the run failed outright and reported that as an error. + /// + /// Read from a clone: [`Agent::run`](crate::Agent::run) takes its options by + /// value, so a caller that wants this keeps a `clone()` of what it passed + /// in. The slot is shared behind an `Arc`, exactly as the counter behind + /// [`reported_tokens`](Self::reported_tokens) is. + pub fn ended_early(&self) -> Option { + self.early_end.get().copied() + } + + /// Records why the runner is ending this run early, keeping the first answer + /// if one is already there. See [`early_end`](Self::early_end) for why the + /// first is the one that stays true. + pub(crate) fn record_early_end(&self, end: EarlyEnd) { + let _ = self.early_end.set(end); + } + + /// Whether cumulative reported usage has reached or exceeded + /// [`token_budget`](Self::token_budget). `false` when no bound is set. + pub(crate) fn token_budget_exceeded(&self) -> bool { + self.token_budget + .is_some_and(|budget| self.reported_tokens() >= budget) + } + + pub(crate) fn check_limits(&self) -> Result<(), RuntimeError> { + if self + .cancellation + .as_ref() + .is_some_and(CancellationToken::is_cancelled) + { + return Err(RuntimeError::Cancelled); + } + + if self + .deadline + .is_some_and(|deadline| SystemTime::now() >= deadline) + { + return Err(RuntimeError::DeadlineExceeded); + } + + Ok(()) + } + + /// Whether a graceful stop has been requested via [`stop`](Self::stop). The + /// runner checks this at each round boundary, where the transcript is at a + /// consistent point, and ends the run successfully when it is set. + pub(crate) fn stop_requested(&self) -> bool { + self.stop + .as_ref() + .is_some_and(CancellationToken::is_cancelled) + } + + pub(crate) fn tool_budget(&self) -> usize { + self.tool_budget.unwrap_or(usize::MAX) + } + + pub(crate) fn model_budget(&self) -> usize { + self.model_budget.unwrap_or(usize::MAX) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + /// What the runner waited before each retry before this type existed: + /// 500 ms doubling to a 5 s ceiling. Pinned so a future edit to the + /// defaults has to be a deliberate one. + #[test] + fn the_default_schedule_is_the_one_mentra_has_always_used() { + let retry = ProviderRetry::default(); + + let delays: Vec = (1..=8) + .map(|attempt| retry.scheduled_delay(attempt)) + .collect(); + + assert_eq!( + delays, + vec![ + Duration::from_millis(500), + Duration::from_secs(1), + Duration::from_secs(2), + Duration::from_secs(4), + Duration::from_secs(5), + Duration::from_secs(5), + Duration::from_secs(5), + Duration::from_secs(5), + ] + ); + } + + #[test] + fn a_host_schedule_doubles_from_its_own_base_to_its_own_ceiling() { + let retry = ProviderRetry { + base_delay: Duration::from_secs(2), + max_delay: Duration::from_secs(10), + ..ProviderRetry::default() + }; + + let delays: Vec = (1..=5) + .map(|attempt| retry.scheduled_delay(attempt)) + .collect(); + + assert_eq!( + delays, + vec![ + Duration::from_secs(2), + Duration::from_secs(4), + Duration::from_secs(8), + Duration::from_secs(10), + Duration::from_secs(10), + ] + ); + } + + #[test] + fn a_long_budget_does_not_overflow_the_doubling() { + // The factor is a `u32`; the old code clamped the shift to `usize`'s + // width, so a host generous enough to allow a 64th attempt got a panic + // instead of a delay. Unreachable at a budget of five, reachable the + // moment raising the budget is the supported thing to do. + let retry = ProviderRetry::default(); + + assert_eq!(retry.scheduled_delay(usize::MAX), retry.max_delay); + assert_eq!(retry.scheduled_delay(64), retry.max_delay); + } + + #[test] + fn a_server_that_names_a_longer_wait_gets_it() { + let retry = ProviderRetry::default(); + + // Attempt 1's own schedule is 500 ms; the limit lasts a minute. + assert_eq!( + retry.delay_for(1, Some(Duration::from_secs(45))), + Duration::from_secs(45) + ); + } + + #[test] + fn a_server_that_names_a_shorter_wait_does_not_shorten_the_schedule() { + // `Retry-After` is the provider's floor on when it will answer, not a + // licence to hammer it sooner than the host chose to. + let retry = ProviderRetry { + base_delay: Duration::from_secs(5), + ..ProviderRetry::default() + }; + + assert_eq!( + retry.delay_for(1, Some(Duration::from_secs(1))), + Duration::from_secs(5) + ); + } + + #[test] + fn a_server_cannot_park_the_run_for_an_hour() { + let retry = ProviderRetry::default(); + + assert_eq!( + retry.delay_for(1, Some(Duration::from_secs(3600))), + retry.retry_after_cap, + "the header is clamped before it is considered" + ); + assert_eq!(retry.retry_after_cap, Duration::from_secs(60)); + } + + #[test] + fn the_cap_bounds_the_server_and_never_the_host() { + // A host that chose to wait five minutes between attempts keeps that + // schedule; the cap exists to bound a remote party, not the caller. + let retry = ProviderRetry { + base_delay: Duration::from_secs(300), + max_delay: Duration::from_secs(300), + retry_after_cap: Duration::from_secs(60), + }; + + assert_eq!(retry.delay_for(1, None), Duration::from_secs(300)); + assert_eq!( + retry.delay_for(1, Some(Duration::from_secs(3600))), + Duration::from_secs(300) + ); + } + + #[test] + fn a_silent_provider_leaves_the_schedule_alone() { + let retry = ProviderRetry::default(); + + assert_eq!(retry.delay_for(3, None), retry.scheduled_delay(3)); + } + + #[test] + fn a_default_run_carries_the_default_schedule() { + assert_eq!( + RunOptions::default().provider_retry, + ProviderRetry::default() + ); + assert_eq!(RunOptions::default().retry_budget, 5); + } + + #[test] + fn a_delegated_run_meets_the_same_provider_with_the_same_patience() { + // A child dials the endpoint its parent dialled, and that endpoint's + // rate-limit window did not shorten because the work was delegated. + // Resetting here would hand the subagent the blip-shaped schedule the + // parent had already found too short against the same limiter. + let parent = RunOptions { + retry_budget: 9, + ..RunOptions::default() + } + .with_provider_retry(ProviderRetry { + base_delay: Duration::from_secs(2), + max_delay: Duration::from_secs(30), + ..ProviderRetry::default() + }); + + let child = parent.child(); + + assert_eq!(child.provider_retry, parent.provider_retry); + assert_eq!(child.retry_budget, 9); + assert_eq!( + child.model_budget, None, + "what the child does not inherit is an allowance for its own work" + ); + } + + #[test] + fn a_cancellation_token_shows_whether_it_was_tripped() { + let token = CancellationToken::default(); + assert!(format!("{token:?}").contains("cancelled: false")); + + token.cancel(); + assert!(format!("{token:?}").contains("cancelled: true")); + } + + #[test] + fn a_clone_of_a_runs_options_reads_what_that_run_recorded() { + // The mechanism an embedder depends on: `Agent::run` takes its options + // by value, so the only way to hear back from a run is to have kept a + // clone — which shares the slot, exactly as it shares the counter. + let options = RunOptions { + token_budget: Some(100), + ..RunOptions::default() + }; + let held = options.clone(); + + options.record_early_end(EarlyEnd::TokenBudget); + + assert_eq!(held.ended_early(), Some(EarlyEnd::TokenBudget)); + } + + #[test] + fn a_child_run_records_its_early_end_apart_from_its_parent() { + // A delegated run that ends on the shared budget has ended its own run, + // not its parent's. The parent reaches its own next boundary and records + // there; until it does, claiming it ended early would be a guess — and a + // wrong one for a parent that goes on to finish its work. + let parent = RunOptions { + token_budget: Some(100), + ..RunOptions::default() + }; + let child = parent.child(); + + child.record_early_end(EarlyEnd::TokenBudget); + + assert_eq!(child.ended_early(), Some(EarlyEnd::TokenBudget)); + assert_eq!(parent.ended_early(), None); + child.record_tokens(60); + assert_eq!( + parent.reported_tokens(), + 60, + "what the two do share is the accounting, unchanged" + ); + } + + #[test] + fn the_first_early_end_recorded_is_the_one_that_stays() { + // Both conditions are sticky under a fixed bound, so a handle reused for + // a second turn ends the same way it ended the first. Keeping the first + // answer makes that explicit rather than depending on it. + let options = RunOptions::default(); + + options.record_early_end(EarlyEnd::StopRequested); + options.record_early_end(EarlyEnd::TokenBudget); + + assert_eq!(options.ended_early(), Some(EarlyEnd::StopRequested)); + } +} diff --git a/vendor/mentra/src/runtime/control/sandbox.rs b/vendor/mentra/src/runtime/control/sandbox.rs new file mode 100644 index 0000000..d271d66 --- /dev/null +++ b/vendor/mentra/src/runtime/control/sandbox.rs @@ -0,0 +1,115 @@ +//! Container and sandbox environment detection. +//! +//! Detects whether the runtime is executing inside a container (Docker, +//! Podman, etc.) or other restricted environment, which informs policy +//! decisions about file access and shell execution. + +use std::path::Path; + +/// Describes the detected execution environment. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ExecutionEnvironment { + /// Running directly on the host OS. + Host, + /// Running inside a Docker container. + Docker, + /// Running inside a generic container (cgroup signals). + Container, + /// Running inside a CI environment. + ContinuousIntegration, +} + +impl std::fmt::Display for ExecutionEnvironment { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Host => write!(f, "host"), + Self::Docker => write!(f, "docker"), + Self::Container => write!(f, "container"), + Self::ContinuousIntegration => write!(f, "ci"), + } + } +} + +/// Detect the current execution environment. +/// +/// Uses multiple heuristics to determine if we're in a container or CI: +/// - `/.dockerenv` file presence → Docker +/// - `/run/.containerenv` file presence → Podman/container +/// - `container=` in `/proc/1/environ` → generic container +/// - CI-related environment variables → CI +pub fn detect_environment() -> ExecutionEnvironment { + // Check for CI environment variables first. + if is_ci_environment() { + return ExecutionEnvironment::ContinuousIntegration; + } + + // Docker detection. + if Path::new("/.dockerenv").exists() { + return ExecutionEnvironment::Docker; + } + + // Podman / generic container detection. + if Path::new("/run/.containerenv").exists() { + return ExecutionEnvironment::Container; + } + + // Check cgroup for container signals (Linux only). + #[cfg(target_os = "linux")] + if is_in_container_cgroup() { + return ExecutionEnvironment::Container; + } + + ExecutionEnvironment::Host +} + +/// Returns `true` if common CI environment variables are set. +fn is_ci_environment() -> bool { + // Standard CI indicators. + std::env::var("CI").is_ok() + || std::env::var("GITHUB_ACTIONS").is_ok() + || std::env::var("GITLAB_CI").is_ok() + || std::env::var("JENKINS_HOME").is_ok() + || std::env::var("CIRCLECI").is_ok() + || std::env::var("BUILDKITE").is_ok() + || std::env::var("TRAVIS").is_ok() +} + +/// Check `/proc/1/cgroup` for container indicators (Linux only). +#[cfg(target_os = "linux")] +fn is_in_container_cgroup() -> bool { + let Ok(cgroup) = std::fs::read_to_string("/proc/1/cgroup") else { + return false; + }; + cgroup.contains("docker") + || cgroup.contains("lxc") + || cgroup.contains("containerd") + || cgroup.contains("kubepods") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn detect_environment_returns_valid_variant() { + let env = detect_environment(); + // We can't control the test environment, but the function should + // always return a valid variant without panicking. + let display = env.to_string(); + assert!( + ["host", "docker", "container", "ci"].contains(&display.as_str()), + "unexpected environment: {display}" + ); + } + + #[test] + fn display_formats_correctly() { + assert_eq!(ExecutionEnvironment::Host.to_string(), "host"); + assert_eq!(ExecutionEnvironment::Docker.to_string(), "docker"); + assert_eq!(ExecutionEnvironment::Container.to_string(), "container"); + assert_eq!( + ExecutionEnvironment::ContinuousIntegration.to_string(), + "ci" + ); + } +} diff --git a/vendor/mentra/src/runtime/error.rs b/vendor/mentra/src/runtime/error.rs new file mode 100644 index 0000000..3068623 --- /dev/null +++ b/vendor/mentra/src/runtime/error.rs @@ -0,0 +1,228 @@ +use crate::provider::{ProviderError, ProviderId}; +use crate::runtime::control::is_transient_provider_error; +use thiserror::Error; + +/// Classifies a [`RuntimeError`] by its recoverability. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ErrorCategory { + /// Transient failure that may succeed on retry. + Retryable, + /// Permanent failure that cannot be retried. + Terminal, + /// Operation continued but state may be inconsistent. + Degraded, +} + +/// Errors produced while configuring, running, or recovering Mentra agents. +#[derive(Debug, Error)] +pub enum RuntimeError { + #[error("{message}", message = provider_not_found_message(.0))] + ProviderNotFound(Option), + #[error("provider '{0}' did not return any models")] + NoModelsAvailable(ProviderId), + #[error("failed to send provider request: {0}")] + FailedToSendRequest(#[source] ProviderError), + #[error("failed to list provider models: {0}")] + FailedToListModels(#[source] ProviderError), + #[error("failed to stream provider response: {0}")] + FailedToStreamResponse(#[source] ProviderError), + #[error("failed to compact history: {0}")] + FailedToCompactHistory(#[source] ProviderError), + #[error("failed to persist transcript: {0}")] + FailedToPersistTranscript(#[source] std::io::Error), + #[error("failed to serialize transcript: {0}")] + FailedToSerializeTranscript(#[source] serde_json::Error), + #[error("failed to load tasks: {0}")] + FailedToLoadTasks(#[source] std::io::Error), + #[error("failed to write tasks: {0}")] + FailedToWriteTasks(#[source] std::io::Error), + #[error("failed to serialize tasks: {0}")] + FailedToSerializeTasks(#[source] serde_json::Error), + #[error("failed to restore tasks: {0}")] + FailedToRestoreTasks(#[source] std::io::Error), + #[error("failed to load team state: {0}")] + FailedToLoadTeam(#[source] std::io::Error), + #[error("failed to write team state: {0}")] + FailedToWriteTeam(#[source] std::io::Error), + #[error("failed to serialize team state: {0}")] + FailedToSerializeTeam(#[source] serde_json::Error), + #[error("failed to deserialize team state: {0}")] + FailedToDeserializeTeam(#[source] serde_json::Error), + #[error("invalid task state: {0}")] + InvalidTask(String), + #[error("invalid team state: {0}")] + InvalidTeam(String), + #[error("operation denied: {0}")] + OperationDenied(String), + #[error("runtime store error: {0}")] + Store(String), + #[error("cannot branch: {0}")] + Branch(#[source] crate::transcript::BranchError), + #[error("lease unavailable: {0}")] + LeaseUnavailable(String), + #[error("operation cancelled")] + Cancelled, + #[error("deadline exceeded")] + DeadlineExceeded, + #[error("tool budget exceeded at {0} call(s)")] + ToolBudgetExceeded(usize), + #[error("model budget exceeded at {0} request(s)")] + ModelBudgetExceeded(usize), + #[error("max rounds exceeded at {0}")] + MaxRoundsExceeded(usize), + #[error("run completed without a final assistant message")] + EmptyAssistantResponse, + #[error("no resumable user turn is available")] + NoResumableTurn, + #[error("invalid tool input for '{name}' ({id}): {source}")] + InvalidToolUseInput { + id: String, + name: String, + #[source] + source: serde_json::Error, + }, + #[error("malformed provider event: {0}")] + MalformedProviderEvent(String), +} + +impl RuntimeError { + /// Returns the [`ErrorCategory`] for this error, classifying it as + /// retryable, terminal, or degraded. + pub fn category(&self) -> ErrorCategory { + match self { + // Provider-backed errors: delegate to transient check. + Self::FailedToSendRequest(source) + | Self::FailedToStreamResponse(source) + | Self::FailedToCompactHistory(source) => { + if is_transient_provider_error(source) { + ErrorCategory::Retryable + } else { + ErrorCategory::Terminal + } + } + + // Listing models is not retryable even if the provider error is transient. + Self::FailedToListModels(_) => ErrorCategory::Terminal, + + // Configuration and logic errors are permanent. + Self::ProviderNotFound(_) + | Self::NoModelsAvailable(_) + | Self::OperationDenied(_) + | Self::Cancelled + | Self::DeadlineExceeded + | Self::ToolBudgetExceeded(_) + | Self::ModelBudgetExceeded(_) + | Self::MaxRoundsExceeded(_) + | Self::EmptyAssistantResponse + | Self::NoResumableTurn + | Self::InvalidToolUseInput { .. } + | Self::MalformedProviderEvent(_) + | Self::InvalidTask(_) + | Self::InvalidTeam(_) + // Naming an entry that is not there is a caller mistake; retrying + // with the same id gets the same answer. + | Self::Branch(_) => ErrorCategory::Terminal, + + // Persistence and store errors: state may be inconsistent. + Self::FailedToPersistTranscript(_) + | Self::FailedToSerializeTranscript(_) + | Self::FailedToLoadTasks(_) + | Self::FailedToWriteTasks(_) + | Self::FailedToSerializeTasks(_) + | Self::FailedToRestoreTasks(_) + | Self::FailedToLoadTeam(_) + | Self::FailedToWriteTeam(_) + | Self::FailedToSerializeTeam(_) + | Self::FailedToDeserializeTeam(_) + | Self::Store(_) + | Self::LeaseUnavailable(_) => ErrorCategory::Degraded, + } + } +} + +fn provider_not_found_message(provider: &Option) -> String { + match provider { + Some(provider) => format!("provider '{provider}' is not registered"), + None => "no providers are registered".to_string(), + } +} + +#[cfg(test)] +mod tests { + use std::error::Error; + + use super::{ErrorCategory, RuntimeError}; + use crate::provider::{ProviderError, ProviderId}; + + #[test] + fn display_mentions_missing_provider() { + let error = RuntimeError::ProviderNotFound(Some(ProviderId::new("custom"))); + + assert_eq!(error.to_string(), "provider 'custom' is not registered"); + } + + #[test] + fn display_mentions_missing_models() { + let error = RuntimeError::NoModelsAvailable(ProviderId::new("custom")); + + assert_eq!( + error.to_string(), + "provider 'custom' did not return any models" + ); + } + + #[test] + fn source_is_exposed_for_wrapped_errors() { + let error = RuntimeError::FailedToSerializeTasks( + serde_json::from_str::("{").expect_err("invalid json"), + ); + + assert!(error.source().is_some()); + } + + #[test] + fn empty_assistant_response_has_clear_display_text() { + assert_eq!( + RuntimeError::EmptyAssistantResponse.to_string(), + "run completed without a final assistant message" + ); + } + + #[test] + fn transient_provider_error_is_retryable() { + let error = RuntimeError::FailedToSendRequest(ProviderError::Retryable { + message: "rate limited".into(), + delay: None, + }); + + assert_eq!(error.category(), ErrorCategory::Retryable); + } + + #[test] + fn permanent_provider_error_is_terminal() { + let error = + RuntimeError::FailedToSendRequest(ProviderError::InvalidRequest("bad body".into())); + + assert_eq!(error.category(), ErrorCategory::Terminal); + } + + #[test] + fn budget_exceeded_is_terminal() { + assert_eq!( + RuntimeError::ToolBudgetExceeded(100).category(), + ErrorCategory::Terminal, + ); + assert_eq!( + RuntimeError::ModelBudgetExceeded(50).category(), + ErrorCategory::Terminal, + ); + } + + #[test] + fn persistence_io_error_is_degraded() { + let io_err = std::io::Error::new(std::io::ErrorKind::PermissionDenied, "disk full"); + let error = RuntimeError::FailedToPersistTranscript(io_err); + + assert_eq!(error.category(), ErrorCategory::Degraded); + } +} diff --git a/vendor/mentra/src/runtime/handle.rs b/vendor/mentra/src/runtime/handle.rs new file mode 100644 index 0000000..73f868d --- /dev/null +++ b/vendor/mentra/src/runtime/handle.rs @@ -0,0 +1,174 @@ +mod agents; +mod construction; +mod execution; +mod tooling; + +use std::{ + any::{Any, TypeId}, + collections::{BTreeSet, HashMap}, + path::{Path, PathBuf}, + sync::{Arc, Mutex, RwLock}, + time::Duration, +}; + +use tokio::sync::watch; + +use crate::{ + agent::{AgentEventBus, AgentSnapshot}, + background::{BackgroundNotification, BackgroundTaskManager, BackgroundTaskSummary}, + compaction::CompactionEngine, + memory::MemoryEngine, + provider::{Provider, ProviderId, ProviderRegistry}, + runtime::{ + control::{ + AuditHook, CommandOutput, CommandRequest, CommandSpec, LocalRuntimeExecutor, + PreExecutionHooks, RuntimeExecutor, RuntimeHookEvent, RuntimeHooks, RuntimePolicy, + read_limited_file, + }, + error::RuntimeError, + store::{RuntimeStore, SqliteRuntimeStore}, + task::{self, TaskAccess}, + }, + team::{ + TeamDispatch, TeamManager, TeamMemberSummary, TeamMessage, TeamProtocolRequestSummary, + TeamRequestFilter, TeammateHost, + }, + tool::{ExecutableTool, ToolAuthorizer, ToolRegistry}, +}; + +use super::skill::SkillLoader; + +#[derive(Clone)] +pub struct RuntimeHandle { + pub(crate) execution: ExecutionServices, + pub(crate) persistence: PersistenceServices, + pub(crate) collaboration: CollaborationServices, + pub(crate) tooling: ToolingServices, + pub(crate) runtime_intrinsics_enabled: bool, + runtime_instance_id: String, + persisted_runtime_identifier: Arc, + lease_keys: Arc>>, + agent_contexts: Arc>>, + provider_registry: Arc>, +} + +#[derive(Clone)] +pub(crate) struct ExecutionServices { + pub(crate) executor: Arc, + pub(crate) policy: Arc, + pub(crate) tool_authorizer: Option>, + pub(crate) hooks: RuntimeHooks, + pub(crate) pre_hooks: PreExecutionHooks, +} + +#[derive(Clone)] +pub(crate) struct PersistenceServices { + pub(crate) store: Arc, + pub(crate) memory: Arc, + pub(crate) compaction: Arc, +} + +#[derive(Clone)] +pub(crate) struct CollaborationServices { + pub(crate) background_tasks: BackgroundTaskManager, + pub(crate) team: TeamManager, + pub(crate) teammate_host: TeammateHost, +} + +#[derive(Clone)] +pub(crate) struct ToolingServices { + pub(crate) tool_registry: Arc>, + pub(crate) scoped_tools: Arc>>, + pub(crate) skill_loader: Arc>>, + pub(crate) app_contexts: Arc>>>, +} + +#[derive(Clone)] +pub(crate) struct AgentObserver { + pub(crate) events: AgentEventBus, + pub(crate) snapshot_tx: watch::Sender, + pub(crate) snapshot: Arc>, +} + +#[derive(Debug, Clone)] +pub(crate) struct AgentExecutionConfig { + pub(crate) name: String, + pub(crate) team_dir: PathBuf, + pub(crate) tasks_dir: PathBuf, + pub(crate) base_dir: PathBuf, + pub(crate) memory_tool_search_limit: usize, + pub(crate) auto_route_shell: bool, + pub(crate) is_teammate: bool, +} + +impl Drop for RuntimeHandle { + fn drop(&mut self) { + if Arc::strong_count(&self.lease_keys) != 1 { + return; + } + + let lease_keys = { + let lease_keys = self.lease_keys.lock().expect("lease key registry poisoned"); + lease_keys.iter().cloned().collect::>() + }; + + for key in lease_keys { + let _ = self + .persistence + .store + .release_lease(&key, &self.runtime_instance_id); + } + } +} + +impl RuntimeHandle { + pub(crate) fn get_provider(&self, id: Option<&ProviderId>) -> Option> { + self.provider_registry + .read() + .expect("provider registry poisoned") + .get_provider(id) + } + + /// The Responses transport this runtime chose for every request it makes, + /// or `None` when it left the choice to each request's own options. + pub(crate) fn responses_transport(&self) -> Option { + self.provider_registry + .read() + .expect("provider registry poisoned") + .responses_transport() + } + + pub(crate) fn memory_engine(&self) -> Arc { + self.persistence.memory.clone() + } + + pub(crate) fn compaction_engine(&self) -> Arc { + self.persistence.compaction.clone() + } + + pub(crate) fn pre_hooks(&self) -> &PreExecutionHooks { + &self.execution.pre_hooks + } + + pub(crate) fn hooks(&self) -> &RuntimeHooks { + &self.execution.hooks + } + + pub(crate) fn with_provider_registry( + &self, + provider_registry: Arc>, + ) -> Self { + Self { + execution: self.execution.clone(), + persistence: self.persistence.clone(), + collaboration: self.collaboration.clone(), + tooling: self.tooling.clone(), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: self.runtime_instance_id.clone(), + persisted_runtime_identifier: self.persisted_runtime_identifier.clone(), + lease_keys: self.lease_keys.clone(), + agent_contexts: self.agent_contexts.clone(), + provider_registry, + } + } +} diff --git a/vendor/mentra/src/runtime/handle/agents.rs b/vendor/mentra/src/runtime/handle/agents.rs new file mode 100644 index 0000000..e1a0d95 --- /dev/null +++ b/vendor/mentra/src/runtime/handle/agents.rs @@ -0,0 +1,185 @@ +use super::*; +use crate::{ + agent::{AgentEvent, AgentEventBus, AgentSnapshot}, + background::{BackgroundObserverSink, BackgroundRegistration}, + team::{TeamObserverSink, TeamRegistration}, +}; + +struct AgentTeamObserver { + store: Arc, + tasks_dir: PathBuf, + events: AgentEventBus, + snapshot_tx: watch::Sender, + snapshot: Arc>, +} + +impl AgentTeamObserver { + fn new( + store: Arc, + tasks_dir: PathBuf, + observer: &AgentObserver, + ) -> Self { + Self { + store, + tasks_dir, + events: observer.events.clone(), + snapshot_tx: observer.snapshot_tx.clone(), + snapshot: Arc::clone(&observer.snapshot), + } + } +} + +impl TeamObserverSink for AgentTeamObserver { + fn publish_snapshot( + &self, + members: &[crate::team::TeamMemberSummary], + requests: &[crate::team::TeamProtocolRequestSummary], + unread_count: usize, + ) { + let mut snapshot = self.snapshot.lock().expect("agent snapshot poisoned"); + if let Ok(tasks) = self.store.load_tasks(self.tasks_dir.as_path()) { + snapshot.tasks = tasks; + } + snapshot.teammates = members.to_vec(); + snapshot.protocol_requests = requests.to_vec(); + snapshot.pending_team_messages = unread_count; + let next_snapshot = snapshot.clone(); + drop(snapshot); + self.snapshot_tx.send_replace(next_snapshot); + } + + fn publish_event(&self, event: AgentEvent) { + self.events.send(event); + } +} + +struct AgentBackgroundObserver { + background_tasks: crate::background::BackgroundTaskManager, + team: crate::team::TeamManager, + agent_id: String, + team_dir: PathBuf, + agent_name: String, + is_teammate: bool, + snapshot_tx: watch::Sender, + snapshot: Arc>, + events: AgentEventBus, +} + +impl AgentBackgroundObserver { + fn new( + background_tasks: crate::background::BackgroundTaskManager, + team: crate::team::TeamManager, + agent_id: String, + config: &AgentExecutionConfig, + observer: &AgentObserver, + ) -> Self { + Self { + background_tasks, + team, + agent_id, + team_dir: config.team_dir.clone(), + agent_name: config.name.clone(), + is_teammate: config.is_teammate, + snapshot_tx: observer.snapshot_tx.clone(), + snapshot: Arc::clone(&observer.snapshot), + events: observer.events.clone(), + } + } +} + +impl BackgroundObserverSink for AgentBackgroundObserver { + fn publish_snapshot(&self, tasks: &[crate::background::BackgroundTaskSummary]) { + let mut snapshot = self.snapshot.lock().expect("agent snapshot poisoned"); + snapshot.background_tasks = tasks.to_vec(); + let next_snapshot = snapshot.clone(); + drop(snapshot); + self.snapshot_tx.send_replace(next_snapshot); + if self.is_teammate + && self + .background_tasks + .has_pending_notifications(&self.agent_id) + { + let _ = self + .team + .wake_teammate(self.team_dir.as_path(), &self.agent_name); + } + } + + fn publish_event(&self, event: AgentEvent) { + let should_wake_teammate = + self.is_teammate && matches!(event, AgentEvent::BackgroundTaskFinished { .. }); + self.events.send(event); + if should_wake_teammate { + let _ = self + .team + .wake_teammate(self.team_dir.as_path(), &self.agent_name); + } + } +} + +impl RuntimeHandle { + pub fn register_agent( + &self, + agent_id: &str, + agent_name: &str, + config: AgentExecutionConfig, + observer: &AgentObserver, + ) -> Result<(), RuntimeError> { + self.acquire_agent_lease(agent_id)?; + self.collaboration + .background_tasks + .register_agent(BackgroundRegistration { + agent_id: agent_id.to_string(), + observer: Arc::new(AgentBackgroundObserver::new( + self.collaboration.background_tasks.clone(), + self.collaboration.team.clone(), + agent_id.to_string(), + &config, + observer, + )), + }); + self.collaboration.team.register_agent(TeamRegistration { + agent_name: agent_name.to_string(), + team_dir: config.team_dir.clone(), + observer: Arc::new(AgentTeamObserver::new( + self.persistence.store.clone(), + config.tasks_dir.clone(), + observer, + )), + })?; + self.agent_contexts + .write() + .expect("agent context registry poisoned") + .insert(agent_id.to_string(), config); + Ok(()) + } + + pub fn acquire_agent_lease(&self, agent_id: &str) -> Result<(), RuntimeError> { + let key = format!("agent:{agent_id}"); + let acquired = self.persistence.store.acquire_lease( + &key, + &self.runtime_instance_id, + Duration::from_secs(3600), + )?; + if acquired { + self.lease_keys + .lock() + .expect("lease key registry poisoned") + .insert(key); + Ok(()) + } else { + Err(RuntimeError::LeaseUnavailable(format!( + "Agent '{agent_id}' is already leased by another runtime" + ))) + } + } + + pub(crate) fn agent_config(&self, agent_id: &str) -> Result { + self.agent_contexts + .read() + .expect("agent context registry poisoned") + .get(agent_id) + .cloned() + .ok_or_else(|| format!("Unknown agent '{agent_id}'")) + } +} diff --git a/vendor/mentra/src/runtime/handle/construction.rs b/vendor/mentra/src/runtime/handle/construction.rs new file mode 100644 index 0000000..c936931 --- /dev/null +++ b/vendor/mentra/src/runtime/handle/construction.rs @@ -0,0 +1,455 @@ +use super::*; +use crate::background::BackgroundHookSink; +use crate::compaction::StandardCompactionEngine; +use crate::memory::MemoryEngine; + +#[derive(Clone)] +struct RuntimeBackgroundHookSink { + store: Arc, + hooks: RuntimeHooks, +} + +impl BackgroundHookSink for RuntimeBackgroundHookSink { + fn task_started( + &self, + agent_id: &str, + task_id: &str, + command: &str, + cwd: &Path, + ) -> Result<(), RuntimeError> { + self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::BackgroundTaskStarted { + agent_id: agent_id.to_string(), + task_id: task_id.to_string(), + command: command.to_string(), + cwd: cwd.to_path_buf(), + }, + ) + } + + fn task_finished( + &self, + agent_id: &str, + task_id: &str, + status: &str, + ) -> Result<(), RuntimeError> { + self.hooks.emit_runtime( + self.store.as_ref(), + &RuntimeHookEvent::BackgroundTaskFinished { + agent_id: agent_id.to_string(), + task_id: task_id.to_string(), + status: status.to_string(), + }, + ) + } +} + +fn background_hook_sink( + store: Arc, + hooks: RuntimeHooks, +) -> Arc { + Arc::new(RuntimeBackgroundHookSink { store, hooks }) +} + +fn clone_tooling_services(tooling: &ToolingServices) -> ToolingServices { + ToolingServices { + tool_registry: Arc::new(RwLock::new( + tooling + .tool_registry + .read() + .expect("tool registry poisoned") + .clone(), + )), + scoped_tools: Arc::new(RwLock::new( + tooling + .scoped_tools + .read() + .expect("scoped tool registry poisoned") + .clone(), + )), + skill_loader: Arc::new(RwLock::new( + tooling + .skill_loader + .read() + .expect("skill loader poisoned") + .clone(), + )), + app_contexts: tooling.app_contexts.clone(), + } +} + +impl RuntimeHandle { + /// Assembles a handle around the default store, without opening it. + /// + /// A builder may replace the store before it settles, so nothing here may + /// touch the database: constructing a [`SqliteRuntimeStore`] only records a + /// path, and the first `open()` is what creates the directory and runs the + /// schema. Recovery is deferred to + /// [`prepare_recovery`](Self::prepare_recovery), which the builder calls + /// once on whichever store the caller kept. + pub fn new(runtime_intrinsics_enabled: bool) -> Self { + let store: Arc = Arc::new(SqliteRuntimeStore::default()); + let executor: Arc = Arc::new(LocalRuntimeExecutor); + let policy = Arc::new(RuntimePolicy::default()); + let hooks = RuntimeHooks::new().with_hook(AuditHook); + let compaction: Arc = + Arc::new(StandardCompactionEngine); + let runtime_instance_id = format!("runtime-{}", std::process::id()); + let memory = Arc::new(MemoryEngine::new(store.clone(), hooks.clone())); + let mut tool_registry = ToolRegistry::default(); + if runtime_intrinsics_enabled { + crate::runtime::intrinsic::register_tools(&mut tool_registry); + tool_registry.register_builtin_tools(crate::tool::FileToolProfile::default()); + } + Self { + execution: ExecutionServices { + executor: executor.clone(), + policy, + tool_authorizer: None, + hooks: hooks.clone(), + pre_hooks: PreExecutionHooks::new(), + }, + persistence: PersistenceServices { + store: store.clone(), + memory, + compaction, + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + store.clone(), + executor, + background_hook_sink(store.clone(), hooks), + ), + team: TeamManager::new(store), + teammate_host: TeammateHost::new().expect("teammate host"), + }, + tooling: ToolingServices { + tool_registry: Arc::new(RwLock::new(tool_registry)), + scoped_tools: Arc::new(RwLock::new(HashMap::new())), + skill_loader: Arc::new(RwLock::new(None)), + app_contexts: Arc::new(RwLock::new(HashMap::new())), + }, + runtime_intrinsics_enabled, + runtime_instance_id, + persisted_runtime_identifier: Arc::::from("default"), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: Arc::new(RwLock::new(ProviderRegistry::default())), + } + } + + /// Reconciles interrupted state on this handle's store and announces it. + /// + /// The builder calls this once, at the build boundary, because that is the + /// first moment the store is known to be final. Calling it earlier — from + /// [`new`](Self::new) or from [`rebind_store`](Self::rebind_store) — opens + /// a database the caller may be about to discard, and writes a second + /// `RecoveryPrepared` audit row that makes "how many times did this runtime + /// start?" unanswerable from the audit trail. + /// + /// Recovery is best-effort: a store that cannot reconcile its interrupted + /// state does not sink an otherwise usable runtime. + pub fn prepare_recovery(&self) { + let _ = self.persistence.store.prepare_recovery(); + let _ = self.emit_hook(RuntimeHookEvent::RecoveryPrepared { + runtime_instance_id: self.runtime_instance_id.clone(), + }); + } + + /// Returns a handle backed by `store` instead of this one's. + /// + /// The replacement is not prepared here; see + /// [`prepare_recovery`](Self::prepare_recovery) for why that waits for the + /// build boundary. + pub fn rebind_store(&self, store: Arc) -> Self { + Self { + execution: self.execution.clone(), + persistence: PersistenceServices { + store: store.clone(), + memory: Arc::new(MemoryEngine::new( + store.clone(), + self.execution.hooks.clone(), + )), + compaction: self.persistence.compaction.clone(), + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + store.clone(), + self.execution.executor.clone(), + background_hook_sink(store.clone(), self.execution.hooks.clone()), + ), + team: TeamManager::new(store), + teammate_host: self.collaboration.teammate_host.clone(), + }, + tooling: clone_tooling_services(&self.tooling), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: format!("runtime-{}", std::process::id()), + persisted_runtime_identifier: self.persisted_runtime_identifier.clone(), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: self.provider_registry.clone(), + } + } + + pub fn with_executor(&self, executor: Arc) -> Self { + Self { + execution: ExecutionServices { + executor: executor.clone(), + policy: self.execution.policy.clone(), + tool_authorizer: self.execution.tool_authorizer.clone(), + hooks: self.execution.hooks.clone(), + pre_hooks: self.execution.pre_hooks.clone(), + }, + persistence: PersistenceServices { + store: self.persistence.store.clone(), + memory: Arc::new(MemoryEngine::new( + self.persistence.store.clone(), + self.execution.hooks.clone(), + )), + compaction: self.persistence.compaction.clone(), + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + self.persistence.store.clone(), + executor, + background_hook_sink( + self.persistence.store.clone(), + self.execution.hooks.clone(), + ), + ), + team: self.collaboration.team.clone(), + teammate_host: self.collaboration.teammate_host.clone(), + }, + tooling: clone_tooling_services(&self.tooling), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: format!("runtime-{}", std::process::id()), + persisted_runtime_identifier: self.persisted_runtime_identifier.clone(), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: self.provider_registry.clone(), + } + } + + pub fn with_policy(&self, policy: RuntimePolicy) -> Self { + Self { + execution: ExecutionServices { + executor: self.execution.executor.clone(), + policy: Arc::new(policy), + tool_authorizer: self.execution.tool_authorizer.clone(), + hooks: self.execution.hooks.clone(), + pre_hooks: self.execution.pre_hooks.clone(), + }, + persistence: PersistenceServices { + store: self.persistence.store.clone(), + memory: Arc::new(MemoryEngine::new( + self.persistence.store.clone(), + self.execution.hooks.clone(), + )), + compaction: self.persistence.compaction.clone(), + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + self.persistence.store.clone(), + self.execution.executor.clone(), + background_hook_sink( + self.persistence.store.clone(), + self.execution.hooks.clone(), + ), + ), + team: self.collaboration.team.clone(), + teammate_host: self.collaboration.teammate_host.clone(), + }, + tooling: clone_tooling_services(&self.tooling), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: format!("runtime-{}", std::process::id()), + persisted_runtime_identifier: self.persisted_runtime_identifier.clone(), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: self.provider_registry.clone(), + } + } + + pub fn with_hooks(&self, hooks: RuntimeHooks) -> Self { + Self { + execution: ExecutionServices { + executor: self.execution.executor.clone(), + policy: self.execution.policy.clone(), + tool_authorizer: self.execution.tool_authorizer.clone(), + hooks: hooks.clone(), + pre_hooks: self.execution.pre_hooks.clone(), + }, + persistence: PersistenceServices { + store: self.persistence.store.clone(), + memory: Arc::new(MemoryEngine::new( + self.persistence.store.clone(), + hooks.clone(), + )), + compaction: self.persistence.compaction.clone(), + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + self.persistence.store.clone(), + self.execution.executor.clone(), + background_hook_sink(self.persistence.store.clone(), hooks), + ), + team: self.collaboration.team.clone(), + teammate_host: self.collaboration.teammate_host.clone(), + }, + tooling: clone_tooling_services(&self.tooling), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: format!("runtime-{}", std::process::id()), + persisted_runtime_identifier: self.persisted_runtime_identifier.clone(), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: self.provider_registry.clone(), + } + } + + pub fn with_pre_hooks(&self, pre_hooks: PreExecutionHooks) -> Self { + Self { + execution: ExecutionServices { + executor: self.execution.executor.clone(), + policy: self.execution.policy.clone(), + tool_authorizer: self.execution.tool_authorizer.clone(), + hooks: self.execution.hooks.clone(), + pre_hooks, + }, + persistence: PersistenceServices { + store: self.persistence.store.clone(), + memory: Arc::new(MemoryEngine::new( + self.persistence.store.clone(), + self.execution.hooks.clone(), + )), + compaction: self.persistence.compaction.clone(), + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + self.persistence.store.clone(), + self.execution.executor.clone(), + background_hook_sink( + self.persistence.store.clone(), + self.execution.hooks.clone(), + ), + ), + team: self.collaboration.team.clone(), + teammate_host: self.collaboration.teammate_host.clone(), + }, + tooling: clone_tooling_services(&self.tooling), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: format!("runtime-{}", std::process::id()), + persisted_runtime_identifier: self.persisted_runtime_identifier.clone(), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: self.provider_registry.clone(), + } + } + + pub fn with_runtime_identifier(&self, runtime_identifier: impl Into>) -> Self { + Self { + execution: self.execution.clone(), + persistence: PersistenceServices { + store: self.persistence.store.clone(), + memory: Arc::new(MemoryEngine::new( + self.persistence.store.clone(), + self.execution.hooks.clone(), + )), + compaction: self.persistence.compaction.clone(), + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + self.persistence.store.clone(), + self.execution.executor.clone(), + background_hook_sink( + self.persistence.store.clone(), + self.execution.hooks.clone(), + ), + ), + team: self.collaboration.team.clone(), + teammate_host: self.collaboration.teammate_host.clone(), + }, + tooling: clone_tooling_services(&self.tooling), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: format!("runtime-{}", std::process::id()), + persisted_runtime_identifier: runtime_identifier.into(), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: self.provider_registry.clone(), + } + } + + pub fn with_tool_authorizer(&self, tool_authorizer: Arc) -> Self { + Self { + execution: ExecutionServices { + executor: self.execution.executor.clone(), + policy: self.execution.policy.clone(), + tool_authorizer: Some(tool_authorizer), + hooks: self.execution.hooks.clone(), + pre_hooks: self.execution.pre_hooks.clone(), + }, + persistence: PersistenceServices { + store: self.persistence.store.clone(), + memory: Arc::new(MemoryEngine::new( + self.persistence.store.clone(), + self.execution.hooks.clone(), + )), + compaction: self.persistence.compaction.clone(), + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + self.persistence.store.clone(), + self.execution.executor.clone(), + background_hook_sink( + self.persistence.store.clone(), + self.execution.hooks.clone(), + ), + ), + team: self.collaboration.team.clone(), + teammate_host: self.collaboration.teammate_host.clone(), + }, + tooling: clone_tooling_services(&self.tooling), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: format!("runtime-{}", std::process::id()), + persisted_runtime_identifier: self.persisted_runtime_identifier.clone(), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: self.provider_registry.clone(), + } + } + + pub fn with_compaction_engine( + &self, + compaction: Arc, + ) -> Self { + Self { + execution: self.execution.clone(), + persistence: PersistenceServices { + store: self.persistence.store.clone(), + memory: Arc::new(MemoryEngine::new( + self.persistence.store.clone(), + self.execution.hooks.clone(), + )), + compaction, + }, + collaboration: CollaborationServices { + background_tasks: BackgroundTaskManager::new( + self.persistence.store.clone(), + self.execution.executor.clone(), + background_hook_sink( + self.persistence.store.clone(), + self.execution.hooks.clone(), + ), + ), + team: self.collaboration.team.clone(), + teammate_host: self.collaboration.teammate_host.clone(), + }, + tooling: clone_tooling_services(&self.tooling), + runtime_intrinsics_enabled: self.runtime_intrinsics_enabled, + runtime_instance_id: format!("runtime-{}", std::process::id()), + persisted_runtime_identifier: self.persisted_runtime_identifier.clone(), + lease_keys: Arc::new(Mutex::new(BTreeSet::new())), + agent_contexts: Arc::new(RwLock::new(HashMap::new())), + provider_registry: self.provider_registry.clone(), + } + } +} diff --git a/vendor/mentra/src/runtime/handle/execution.rs b/vendor/mentra/src/runtime/handle/execution.rs new file mode 100644 index 0000000..0ada57d --- /dev/null +++ b/vendor/mentra/src/runtime/handle/execution.rs @@ -0,0 +1,575 @@ +use crate::runtime::TaskIntrinsicTool; + +use super::*; + +impl RuntimeHandle { + /// Authorizes, validates and shapes one command into a request. + /// + /// `target` rides through untouched: it names the executor the host wants, + /// and every check above it — working-root authorization, shell + /// validation, the timeout clamp, the output cap, the environment + /// allowlist — applies to a targeted command exactly as it does to a local + /// one. A target chooses where an authorized command runs; it never + /// decides whether it may. + fn build_command_request( + &self, + agent_id: &str, + target: Option, + command: String, + requested_timeout: Option, + cwd: PathBuf, + background: bool, + ) -> Result<(AgentExecutionConfig, CommandRequest), String> { + let config = self.agent_config(agent_id)?; + if let Err(detail) = + self.execution + .policy + .authorize_command_execution(&config.base_dir, &cwd, background) + { + let _ = self.emit_hook(RuntimeHookEvent::AuthorizationDenied { + agent_id: agent_id.to_string(), + action: if background { + "background_command".to_string() + } else { + "shell_command".to_string() + }, + detail: detail.clone(), + }); + return Err(detail); + } + + let validation = self + .execution + .policy + .evaluate_shell_command(&command, &config.base_dir); + if validation.should_emit_hook() { + let detail = validation + .reason() + .map(ToOwned::to_owned) + .unwrap_or_else(|| "Shell command requires validation".to_string()); + let _ = self.emit_hook(RuntimeHookEvent::AuthorizationDenied { + agent_id: agent_id.to_string(), + action: if background { + "background_shell_validation".to_string() + } else { + "shell_validation".to_string() + }, + detail: detail.clone(), + }); + if validation.should_deny() { + return Err(detail); + } + } + + let command_request = CommandRequest { + spec: CommandSpec::Shell { command }, + cwd, + timeout: self.execution.policy.effective_timeout(requested_timeout), + env: self.execution.policy.allowed_environment(), + max_output_bytes_per_stream: self.execution.policy.max_output_bytes_per_stream, + target, + }; + + Ok((config, command_request)) + } + + pub fn start_background_task( + &self, + agent_id: &str, + command: String, + _justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + // Background tasks are untargeted in this release: a task outlives the + // call that started it, and nothing yet reports a remote task's fate + // back to the agent that asked for it. + let (_config, command_request) = + self.build_command_request(agent_id, None, command, requested_timeout, cwd, true)?; + + if let Some(limit) = self.execution.policy.background_task_limit + && self + .collaboration + .background_tasks + .running_task_count(agent_id) + >= limit + { + let detail = format!("Background task limit of {limit} reached"); + let _ = self.emit_hook(RuntimeHookEvent::AuthorizationDenied { + agent_id: agent_id.to_string(), + action: "background_limit".to_string(), + detail: detail.clone(), + }); + return Err(detail); + } + + self.collaboration + .background_tasks + .start_task(agent_id, command_request) + } + + pub fn check_background_task( + &self, + agent_id: &str, + task_id: Option<&str>, + ) -> Result { + self.collaboration + .background_tasks + .check_task(agent_id, task_id) + } + + pub fn drain_background_notifications(&self, agent_id: &str) -> Vec { + self.collaboration + .background_tasks + .drain_notifications(agent_id) + } + + pub fn has_deliverable_background_notifications(&self, agent_id: &str) -> bool { + self.collaboration + .background_tasks + .has_deliverable_notifications(agent_id) + } + + pub fn requeue_background_notifications( + &self, + agent_id: &str, + notifications: Vec, + ) { + self.collaboration + .background_tasks + .requeue_notifications(agent_id, notifications); + } + + pub fn acknowledge_background_notifications(&self, agent_id: &str) { + self.collaboration + .background_tasks + .acknowledge_notifications(agent_id); + } + + pub fn spawn_teammate_actor( + &self, + team_dir: &Path, + teammate_name: &str, + agent: std::sync::Arc>, + ) -> Result { + Ok(self.collaboration.teammate_host.spawn_teammate( + self.collaboration.team.clone(), + team_dir.to_path_buf(), + teammate_name.to_string(), + agent, + )) + } + + pub fn register_teammate( + &self, + team_dir: &Path, + summary: TeamMemberSummary, + actor: crate::team::TeammateActorHandle, + ) -> Result { + self.collaboration + .team + .spawn_teammate(team_dir, summary, actor) + } + + pub fn wake_teammate(&self, team_dir: &Path, teammate_name: &str) -> Result<(), RuntimeError> { + self.collaboration + .team + .wake_teammate(team_dir, teammate_name) + } + + pub fn send_team_message( + &self, + team_dir: &Path, + sender: &str, + to: &str, + content: String, + ) -> Result { + self.collaboration + .team + .send_message(team_dir, sender, to, content) + } + + pub fn broadcast_team_message( + &self, + team_dir: &Path, + sender: &str, + content: String, + ) -> Result, RuntimeError> { + self.collaboration + .team + .broadcast_message(team_dir, sender, content) + } + + pub fn read_team_inbox( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result, RuntimeError> { + self.collaboration.team.read_inbox(team_dir, agent_name) + } + + pub fn requeue_team_messages( + &self, + team_dir: &Path, + agent_name: &str, + messages: Vec, + ) -> Result<(), RuntimeError> { + self.collaboration + .team + .requeue_messages(team_dir, agent_name, messages) + } + + pub fn acknowledge_team_messages( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result<(), RuntimeError> { + self.collaboration + .team + .acknowledge_messages(team_dir, agent_name) + } + + pub fn create_team_request( + &self, + team_dir: &Path, + sender: &str, + to: &str, + protocol: String, + content: String, + ) -> Result { + self.collaboration + .team + .create_request(team_dir, sender, to, protocol, content) + } + + pub fn resolve_team_request( + &self, + team_dir: &Path, + responder: &str, + request_id: &str, + approve: bool, + reason: Option, + ) -> Result { + self.collaboration + .team + .resolve_request(team_dir, responder, request_id, approve, reason) + } + + pub fn list_team_requests( + &self, + team_dir: &Path, + agent_name: &str, + filter: TeamRequestFilter, + ) -> Result, RuntimeError> { + self.collaboration + .team + .list_requests(team_dir, agent_name, filter) + } + + pub fn execute_task_mutation( + &self, + tool: &TaskIntrinsicTool, + input: serde_json::Value, + dir: &Path, + access: TaskAccess<'_>, + ) -> Result { + task::execute_with_store(self.persistence.store.as_ref(), tool, input, dir, access) + } + + /// Runs one command on the local executor. + pub async fn execute_shell_command( + &self, + agent_id: &str, + command: String, + justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + self.execute_shell_command_on( + agent_id, + None, + command, + justification, + requested_timeout, + cwd, + ) + .await + } + + /// Runs one command on the executor the host named. + /// + /// `target` is passed to the installed [`RuntimeExecutor`] on the request + /// and is not interpreted here: which names exist, and what each one + /// reaches, is the host's business. Every guard around the command is the + /// same one an untargeted call gets. `None` means the local executor; + /// the builtin [`LocalRuntimeExecutor`] refuses any other name rather than + /// running a command that was addressed elsewhere. + pub async fn execute_shell_command_on( + &self, + agent_id: &str, + target: Option, + command: String, + _justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + let (_config, command_request) = + self.build_command_request(agent_id, target, command, requested_timeout, cwd, false)?; + + self.execution.executor.run(command_request).await + } + + pub async fn read_file( + &self, + agent_id: &str, + path: &str, + max_lines: Option, + ) -> Result { + let config = self.agent_config(agent_id)?; + let resolved = match self + .execution + .policy + .authorize_file_read(&config.base_dir, Path::new(path)) + { + Ok(path) => path, + Err(detail) => { + let _ = self.emit_hook(RuntimeHookEvent::AuthorizationDenied { + agent_id: agent_id.to_string(), + action: "read_file".to_string(), + detail: detail.clone(), + }); + return Err(detail); + } + }; + + read_limited_file(&resolved, max_lines).await + } + + pub fn resolve_working_directory( + &self, + agent_id: &str, + explicit_directory: Option<&str>, + ) -> Result { + let config = self.agent_config(agent_id)?; + + if let Some(directory) = explicit_directory { + return Ok(resolve_path(&config.base_dir, directory)); + } + + if !config.auto_route_shell { + return Ok(config.base_dir); + } + + let tasks = self + .persistence + .store + .load_tasks(&config.tasks_dir) + .map_err(|error| error.to_string())?; + let owned = tasks + .into_iter() + .filter(|task| { + config.is_teammate + && task.owner == config.name + && !matches!(task.status, crate::runtime::TaskStatus::Completed) + }) + .collect::>(); + + let directories = owned + .iter() + .filter_map(|task| task.working_directory.as_deref()) + .map(|path| resolve_path(&config.base_dir, path)) + .collect::>(); + + if directories.is_empty() { + return Ok(config.base_dir); + } + + if directories.len() > 1 { + return Err( + "Multiple owned task directories are active. Pass workingDirectory explicitly." + .to_string(), + ); + } + + Ok(directories.into_iter().next().expect("one directory")) + } + + pub fn default_working_directory(&self, agent_id: &str) -> PathBuf { + self.agent_contexts + .read() + .expect("agent context registry poisoned") + .get(agent_id) + .map(|config| config.base_dir.clone()) + .unwrap_or_else(|| PathBuf::from(".")) + } + + pub(crate) fn shell_validation( + &self, + agent_id: &str, + command: &str, + ) -> Result { + let config = self.agent_config(agent_id)?; + Ok(self + .execution + .policy + .evaluate_shell_command(command, &config.base_dir)) + } + + pub fn emit_hook(&self, event: RuntimeHookEvent) -> Result<(), RuntimeError> { + self.execution + .hooks + .emit_runtime(self.persistence.store.as_ref(), &event) + } +} + +fn resolve_path(base_dir: &Path, path: &str) -> PathBuf { + let candidate = PathBuf::from(path); + if candidate.is_absolute() { + candidate + } else { + base_dir.join(candidate) + } +} + +#[cfg(test)] +mod tests { + use async_trait::async_trait; + + use super::*; + use crate::runtime::{ + VolatileRuntimeStore, + control::{CommandOutput, LocalRuntimeExecutor, RuntimeExecutor, RuntimePolicy}, + }; + + const AGENT_ID: &str = "agent-1"; + + /// Records what the handle handed it and answers without running anything, + /// so a test can read the request the routing layer actually produced. + #[derive(Default)] + struct RecordingExecutor { + requests: Mutex>, + } + + impl RecordingExecutor { + fn last_target(&self) -> Option { + self.requests + .lock() + .expect("recorded requests poisoned") + .last() + .expect("one recorded request") + .target + .clone() + } + } + + #[async_trait] + impl RuntimeExecutor for RecordingExecutor { + async fn run(&self, request: CommandRequest) -> Result { + self.requests + .lock() + .expect("recorded requests poisoned") + .push(request); + Ok(CommandOutput { + stdout: "recorded".to_string(), + stderr: String::new(), + success: true, + status_code: Some(0), + timed_out: false, + stdout_truncated: false, + stderr_truncated: false, + }) + } + } + + /// A handle wired to `executor`, with one agent registered and a policy + /// that permits shell commands. The store is volatile so nothing here + /// touches the machine-wide database. + fn handle_with(executor: Arc) -> RuntimeHandle { + let handle = RuntimeHandle::new(false) + .rebind_store(Arc::new(VolatileRuntimeStore::new())) + .with_policy(RuntimePolicy::permissive()) + .with_executor(executor); + let base_dir = std::env::temp_dir(); + handle + .agent_contexts + .write() + .expect("agent context registry poisoned") + .insert( + AGENT_ID.to_string(), + AgentExecutionConfig { + name: "agent".to_string(), + team_dir: base_dir.clone(), + tasks_dir: base_dir.clone(), + base_dir, + memory_tool_search_limit: 5, + auto_route_shell: false, + is_teammate: false, + }, + ); + handle + } + + #[tokio::test] + async fn a_named_target_reaches_the_executor() { + let executor = Arc::new(RecordingExecutor::default()); + let handle = handle_with(executor.clone()); + + handle + .execute_shell_command_on( + AGENT_ID, + Some("x".to_string()), + "true".to_string(), + None, + None, + std::env::temp_dir(), + ) + .await + .expect("the stub executor answers"); + + assert_eq!(executor.last_target(), Some("x".to_string())); + } + + #[tokio::test] + async fn an_untargeted_command_reaches_the_executor_with_no_target() { + let executor = Arc::new(RecordingExecutor::default()); + let handle = handle_with(executor.clone()); + + handle + .execute_shell_command( + AGENT_ID, + "true".to_string(), + None, + None, + std::env::temp_dir(), + ) + .await + .expect("the stub executor answers"); + + assert_eq!(executor.last_target(), None); + } + + /// The refusal has to come from the executor, not from a local run that + /// happened to succeed: a command addressed to a host that this runtime + /// cannot reach must fail loudly rather than execute here. + #[tokio::test] + async fn the_local_executor_refuses_a_target_it_does_not_serve() { + let handle = handle_with(Arc::new(LocalRuntimeExecutor)); + + let error = handle + .execute_shell_command_on( + AGENT_ID, + Some("mac".to_string()), + "true".to_string(), + None, + None, + std::env::temp_dir(), + ) + .await + .expect_err("a targeted command must not run on the local executor"); + + assert_eq!( + error, + "no executor serves target `mac`; the local executor only runs untargeted commands" + ); + } +} diff --git a/vendor/mentra/src/runtime/handle/tooling.rs b/vendor/mentra/src/runtime/handle/tooling.rs new file mode 100644 index 0000000..6c68868 --- /dev/null +++ b/vendor/mentra/src/runtime/handle/tooling.rs @@ -0,0 +1,192 @@ +use super::*; + +impl RuntimeHandle { + pub fn configure_file_tools(&self, profile: crate::tool::FileToolProfile) { + self.tooling + .tool_registry + .write() + .expect("tool registry poisoned") + .configure_file_tools(profile); + } + + pub fn register_app_context(&self, context: Arc) { + self.tooling + .app_contexts + .write() + .expect("app context registry poisoned") + .insert(context.as_ref().type_id(), context); + } + + pub fn app_context(&self) -> Result, String> + where + T: Any + Send + Sync + 'static, + { + let context = self + .tooling + .app_contexts + .read() + .expect("app context registry poisoned") + .get(&TypeId::of::()) + .cloned() + .ok_or_else(|| { + format!( + "App context '{}' is not registered on this runtime", + std::any::type_name::() + ) + })?; + + Arc::downcast::(context).map_err(|_| { + format!( + "App context '{}' was registered with an incompatible type", + std::any::type_name::() + ) + }) + } + + pub fn register_tool(&self, tool: T) + where + T: ExecutableTool + 'static, + { + self.tooling + .tool_registry + .write() + .expect("tool registry poisoned") + .register_tool(tool); + } + + pub(crate) fn register_scoped_tool(&self, agent_id: &str, tool: T) + where + T: ExecutableTool + 'static, + { + let name = tool.descriptor().provider.name; + self.tooling + .scoped_tools + .write() + .expect("scoped tool registry poisoned") + .insert(name, agent_id.to_string()); + self.register_tool(tool); + } + + pub(crate) fn unregister_scoped_tool(&self, agent_id: &str, name: &str) { + let owner_matches = self + .tooling + .scoped_tools + .read() + .expect("scoped tool registry poisoned") + .get(name) + .is_some_and(|owner| owner == agent_id); + if !owner_matches { + return; + } + + self.tooling + .tool_registry + .write() + .expect("tool registry poisoned") + .unregister_tool(name); + self.tooling + .scoped_tools + .write() + .expect("scoped tool registry poisoned") + .remove(name); + } + + pub(crate) fn tool_is_visible_to_agent(&self, name: &str, agent_id: &str) -> bool { + self.tooling + .scoped_tools + .read() + .expect("scoped tool registry poisoned") + .get(name) + .is_none_or(|owner| owner == agent_id) + } + + /// Adds a skill root, keeping any name already registered. + /// + /// Additive rather than replacing: registering a second root used to + /// discard the first silently, which made "project skills layered over + /// personal ones" impossible to express. + pub fn register_skill_loader(&self, loader: SkillLoader) { + let mut slot = self + .tooling + .skill_loader + .write() + .expect("skill loader poisoned"); + match slot.as_mut() { + Some(existing) => existing.merge_weaker(loader), + None => *slot = Some(loader), + } + drop(slot); + + self.tooling + .tool_registry + .write() + .expect("tool registry poisoned") + .register_skill_tool(); + } + + /// Every loaded skill, name-ordered, without bodies. + pub fn skills(&self) -> Vec { + self.tooling + .skill_loader + .read() + .expect("skill loader poisoned") + .as_ref() + .map(SkillLoader::infos) + .unwrap_or_default() + } + + pub fn tools(&self) -> Arc<[crate::tool::ProviderToolSpec]> { + self.tooling + .tool_registry + .read() + .expect("tool registry poisoned") + .tools() + } + + pub fn store(&self) -> Arc { + self.persistence.store.clone() + } + + pub fn persisted_runtime_identifier(&self) -> &str { + &self.persisted_runtime_identifier + } + + pub fn skill_descriptions(&self) -> Option { + self.tooling + .skill_loader + .read() + .expect("skill loader poisoned") + .as_ref() + .map(SkillLoader::get_descriptions) + .filter(|descriptions| !descriptions.is_empty()) + } + + pub fn load_skill(&self, name: &str) -> Result { + let skills = self + .tooling + .skill_loader + .read() + .expect("skill loader poisoned"); + let Some(loader) = skills.as_ref() else { + return Err("Skill loader is not available".to_string()); + }; + + loader.get_content(name) + } + + pub fn get_tool(&self, name: &str) -> Option> { + self.tooling + .tool_registry + .read() + .expect("tool registry poisoned") + .get_tool(name) + } + + pub fn get_tool_descriptor(&self, name: &str) -> Option { + self.tooling + .tool_registry + .read() + .expect("tool registry poisoned") + .get_tool_descriptor(name) + } +} diff --git a/vendor/mentra/src/runtime/hybrid_store.rs b/vendor/mentra/src/runtime/hybrid_store.rs new file mode 100644 index 0000000..9b34f76 --- /dev/null +++ b/vendor/mentra/src/runtime/hybrid_store.rs @@ -0,0 +1,438 @@ +use std::{ + path::{Path, PathBuf}, + time::Duration, +}; + +use crate::{ + background::{BackgroundNotification, BackgroundStore, BackgroundTaskSummary}, + memory::{ + MemoryCursor, MemoryListPage, MemoryListRequest, MemoryRecord, MemorySearchRequest, + MemoryStore, SqliteHybridMemoryStore, + }, + runtime::{ + AgentStore, AuditStore, LeaseStore, LoadedAgentState, PermissionRuleStore, + PersistedAgentRecord, RunStore, SqliteRuntimeStore, TaskStateSnapshot, TaskStore, + }, + session::permission::RememberedRule, + team::{TeamMemberSummary, TeamMessage, TeamProtocolRequestSummary, TeamStore}, +}; + +use super::{RuntimeError, TaskItem}; + +#[derive(Clone)] +/// Runtime store that keeps SQLite for runtime state and uses the hybrid memory +/// store for long-term memory records. +pub struct HybridRuntimeStore { + inner: SqliteRuntimeStore, + memory: SqliteHybridMemoryStore, +} + +impl Default for HybridRuntimeStore { + fn default() -> Self { + Self::new(SqliteRuntimeStore::default_path()) + } +} + +impl HybridRuntimeStore { + /// Creates a hybrid store that colocates runtime state with a derived memory database. + pub fn new(runtime_path: impl Into) -> Self { + let runtime_path = runtime_path.into(); + let memory_path = derive_memory_path(runtime_path.as_path()); + Self { + inner: SqliteRuntimeStore::new(runtime_path), + memory: SqliteHybridMemoryStore::new(memory_path), + } + } + + /// Creates a hybrid store with explicit runtime and memory database paths. + pub fn with_memory_path( + runtime_path: impl Into, + memory_path: impl Into, + ) -> Self { + Self { + inner: SqliteRuntimeStore::new(runtime_path), + memory: SqliteHybridMemoryStore::new(memory_path), + } + } + + /// Creates a hybrid store using the default runtime-scoped SQLite path layout. + pub fn for_runtime_identifier(runtime_identifier: &str) -> Self { + Self::new(SqliteRuntimeStore::path_for_runtime_identifier( + runtime_identifier, + )) + } + + pub fn runtime_path(&self) -> &Path { + self.inner.path() + } + + pub fn memory_path(&self) -> &Path { + self.memory.path() + } +} + +impl TeamStore for HybridRuntimeStore { + fn unread_team_count(&self, team_dir: &Path, agent_name: &str) -> Result { + self.inner.unread_team_count(team_dir, agent_name) + } + + fn load_team_members(&self, team_dir: &Path) -> Result, RuntimeError> { + self.inner.load_team_members(team_dir) + } + + fn upsert_team_member( + &self, + team_dir: &Path, + summary: &TeamMemberSummary, + ) -> Result<(), RuntimeError> { + self.inner.upsert_team_member(team_dir, summary) + } + + fn read_team_inbox( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result, RuntimeError> { + self.inner.read_team_inbox(team_dir, agent_name) + } + + fn ack_team_inbox(&self, team_dir: &Path, agent_name: &str) -> Result<(), RuntimeError> { + self.inner.ack_team_inbox(team_dir, agent_name) + } + + fn requeue_team_inbox(&self, team_dir: &Path, agent_name: &str) -> Result<(), RuntimeError> { + self.inner.requeue_team_inbox(team_dir, agent_name) + } + + fn append_team_message( + &self, + team_dir: &Path, + recipient: &str, + message: &TeamMessage, + ) -> Result<(), RuntimeError> { + self.inner.append_team_message(team_dir, recipient, message) + } + + fn load_team_requests( + &self, + team_dir: &Path, + ) -> Result, RuntimeError> { + self.inner.load_team_requests(team_dir) + } + + fn upsert_team_request( + &self, + team_dir: &Path, + request: &TeamProtocolRequestSummary, + ) -> Result<(), RuntimeError> { + self.inner.upsert_team_request(team_dir, request) + } + + fn list_team_agent_names(&self, team_dir: &Path) -> Result, RuntimeError> { + self.inner.list_team_agent_names(team_dir) + } +} + +impl BackgroundStore for HybridRuntimeStore { + fn load_background_tasks( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + self.inner.load_background_tasks(agent_id) + } + + fn upsert_background_task( + &self, + agent_id: &str, + task: &BackgroundTaskSummary, + notification_state: i64, + ) -> Result<(), RuntimeError> { + self.inner + .upsert_background_task(agent_id, task, notification_state) + } + + fn drain_background_notifications( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + self.inner.drain_background_notifications(agent_id) + } + + fn has_pending_background_notifications(&self, agent_id: &str) -> Result { + self.inner.has_pending_background_notifications(agent_id) + } + + fn has_deliverable_background_notifications( + &self, + agent_id: &str, + ) -> Result { + self.inner + .has_deliverable_background_notifications(agent_id) + } + + fn ack_background_notifications(&self, agent_id: &str) -> Result<(), RuntimeError> { + self.inner.ack_background_notifications(agent_id) + } + + fn requeue_background_notifications(&self, agent_id: &str) -> Result<(), RuntimeError> { + self.inner.requeue_background_notifications(agent_id) + } +} + +impl MemoryStore for HybridRuntimeStore { + fn upsert_records(&self, records: &[MemoryRecord]) -> Result<(), RuntimeError> { + self.memory.upsert_records(records) + } + + fn search_records_with_options( + &self, + request: &MemorySearchRequest, + ) -> Result, RuntimeError> { + self.memory.search_records_with_options(request) + } + + fn search_records( + &self, + agent_id: &str, + query: &str, + limit: usize, + ) -> Result, RuntimeError> { + self.memory.search_records(agent_id, query, limit) + } + + fn list_records(&self, request: &MemoryListRequest) -> Result { + self.memory.list_records(request) + } + + fn get_record( + &self, + agent_id: &str, + record_id: &str, + ) -> Result, RuntimeError> { + self.memory.get_record(agent_id, record_id) + } + + fn count_records(&self, agent_id: &str) -> Result { + self.memory.count_records(agent_id) + } + + fn delete_records(&self, record_ids: &[String]) -> Result<(), RuntimeError> { + self.memory.delete_records(record_ids) + } + + fn tombstone_records( + &self, + agent_id: &str, + record_ids: &[String], + ) -> Result { + self.memory.tombstone_records(agent_id, record_ids) + } + + fn load_agent_memory_cursor( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + self.memory.load_agent_memory_cursor(agent_id) + } + + fn save_agent_memory_cursor( + &self, + agent_id: &str, + cursor: &MemoryCursor, + ) -> Result<(), RuntimeError> { + self.memory.save_agent_memory_cursor(agent_id, cursor) + } +} + +impl AgentStore for HybridRuntimeStore { + fn prepare_recovery(&self) -> Result<(), RuntimeError> { + self.inner.prepare_recovery() + } + + fn create_agent( + &self, + record: &PersistedAgentRecord, + memory: &crate::memory::journal::AgentMemoryState, + ) -> Result<(), RuntimeError> { + self.inner.create_agent(record, memory) + } + + fn save_agent_record(&self, record: &PersistedAgentRecord) -> Result<(), RuntimeError> { + self.inner.save_agent_record(record) + } + + fn save_agent_memory( + &self, + agent_id: &str, + memory: &crate::memory::journal::AgentMemoryState, + ) -> Result<(), RuntimeError> { + self.inner.save_agent_memory(agent_id, memory) + } + + fn load_agent(&self, agent_id: &str) -> Result, RuntimeError> { + self.inner.load_agent(agent_id) + } + + fn list_agents(&self) -> Result, RuntimeError> { + self.inner.list_agents() + } + + fn list_agents_by_runtime( + &self, + runtime_identifier: &str, + ) -> Result, RuntimeError> { + self.inner.list_agents_by_runtime(runtime_identifier) + } +} + +impl RunStore for HybridRuntimeStore { + fn start_run(&self, agent_id: &str) -> Result { + self.inner.start_run(agent_id) + } + + fn update_run_state( + &self, + run_id: &str, + state: &str, + error: Option<&str>, + ) -> Result<(), RuntimeError> { + self.inner.update_run_state(run_id, state, error) + } + + fn finish_run(&self, run_id: &str) -> Result<(), RuntimeError> { + self.inner.finish_run(run_id) + } + + fn fail_run(&self, run_id: &str, error: &str) -> Result<(), RuntimeError> { + self.inner.fail_run(run_id, error) + } +} + +impl TaskStore for HybridRuntimeStore { + fn load_tasks(&self, namespace: &Path) -> Result, RuntimeError> { + self.inner.load_tasks(namespace) + } + + fn capture_tasks(&self, namespace: &Path) -> Result { + self.inner.capture_tasks(namespace) + } + + fn restore_tasks( + &self, + namespace: &Path, + snapshot: &TaskStateSnapshot, + ) -> Result<(), RuntimeError> { + self.inner.restore_tasks(namespace, snapshot) + } + + fn replace_tasks(&self, namespace: &Path, tasks: &[TaskItem]) -> Result<(), RuntimeError> { + self.inner.replace_tasks(namespace, tasks) + } + + fn mutate( + &self, + namespace: &Path, + mutation: &mut dyn FnMut(&mut Vec) -> Result<(), RuntimeError>, + ) -> Result<(), RuntimeError> { + self.inner.mutate(namespace, mutation) + } +} + +impl AuditStore for HybridRuntimeStore { + fn record_audit_event( + &self, + scope: &str, + event_type: &str, + payload: serde_json::Value, + ) -> Result<(), RuntimeError> { + self.inner.record_audit_event(scope, event_type, payload) + } +} + +impl LeaseStore for HybridRuntimeStore { + fn acquire_lease(&self, key: &str, owner: &str, ttl: Duration) -> Result { + self.inner.acquire_lease(key, owner, ttl) + } + + fn release_lease(&self, key: &str, owner: &str) -> Result<(), RuntimeError> { + self.inner.release_lease(key, owner) + } +} + +impl PermissionRuleStore for HybridRuntimeStore { + fn save_rules( + &self, + session_id: &str, + project_id: Option<&str>, + rules: &[RememberedRule], + ) -> Result<(), RuntimeError> { + self.inner.save_rules(session_id, project_id, rules) + } + + fn load_rules( + &self, + session_id: &str, + project_id: Option<&str>, + ) -> Result, RuntimeError> { + self.inner.load_rules(session_id, project_id) + } + + fn clear_rules(&self, session_id: &str) -> Result<(), RuntimeError> { + self.inner.clear_rules(session_id) + } +} + +fn derive_memory_path(runtime_path: &Path) -> PathBuf { + let stem = runtime_path + .file_stem() + .and_then(|value| value.to_str()) + .unwrap_or("runtime"); + let file_name = format!("{stem}-memory.sqlite"); + runtime_path.with_file_name(file_name) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + agent::{AgentConfig, AgentStatus}, + runtime::AgentStore, + }; + + #[test] + fn wrapper_store_delegates_non_memory_runtime_operations() { + let base = std::env::temp_dir().join(format!( + "mentra-hybrid-runtime-{}.sqlite", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("time") + .as_nanos() + )); + let store = HybridRuntimeStore::new(base); + let record = PersistedAgentRecord { + id: "agent-1".to_string(), + runtime_identifier: "default".to_string(), + name: "agent".to_string(), + model: "model".to_string(), + provider_id: "anthropic".into(), + config: AgentConfig::default(), + hidden_tools: Default::default(), + max_rounds: None, + teammate_identity: None, + rounds_since_task: 0, + idle_requested: false, + status: AgentStatus::Idle, + subagents: Vec::new(), + }; + store + .create_agent( + &record, + &crate::memory::journal::AgentMemoryState::default(), + ) + .expect("create agent"); + + let loaded = store.load_agent("agent-1").expect("load agent"); + assert!(loaded.is_some()); + assert_ne!(store.runtime_path(), store.memory_path()); + } +} diff --git a/vendor/mentra/src/runtime/intrinsic.rs b/vendor/mentra/src/runtime/intrinsic.rs new file mode 100644 index 0000000..7f571d0 --- /dev/null +++ b/vendor/mentra/src/runtime/intrinsic.rs @@ -0,0 +1,52 @@ +#[path = "intrinsic/descriptor.rs"] +mod descriptor; +#[path = "intrinsic/execute.rs"] +mod execute; + +use async_trait::async_trait; +use strum::{Display, VariantArray}; + +use crate::tool::{ + ParallelToolContext, RuntimeToolDescriptor, ToolContext, ToolDefinition, ToolExecutor, + ToolResult, +}; + +pub(crate) fn register_tools(registry: &mut crate::tool::ToolRegistry) { + RuntimeIntrinsicTool::VARIANTS + .iter() + .for_each(|tool| registry.register_tool(*tool)); + crate::runtime::task::TaskIntrinsicTool::VARIANTS + .iter() + .for_each(|tool| registry.register_tool(*tool)); + crate::team::TeamIntrinsicTool::VARIANTS + .iter() + .for_each(|tool| registry.register_tool(*tool)); +} + +#[derive(Display, Copy, Clone, VariantArray)] +#[strum(serialize_all = "snake_case")] +pub(crate) enum RuntimeIntrinsicTool { + Compact, + Idle, + MemoryForget, + MemoryPin, + MemorySearch, + Task, +} + +impl ToolDefinition for RuntimeIntrinsicTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + descriptor::runtime_intrinsic_descriptor(*self) + } +} + +#[async_trait] +impl ToolExecutor for RuntimeIntrinsicTool { + async fn execute(&self, ctx: ParallelToolContext, input: serde_json::Value) -> ToolResult { + execute::execute_parallel(*self, ctx, input).await + } + + async fn execute_mut(&self, ctx: ToolContext<'_>, input: serde_json::Value) -> ToolResult { + execute::execute_mut(*self, ctx, input).await + } +} diff --git a/vendor/mentra/src/runtime/intrinsic/descriptor.rs b/vendor/mentra/src/runtime/intrinsic/descriptor.rs new file mode 100644 index 0000000..1425662 --- /dev/null +++ b/vendor/mentra/src/runtime/intrinsic/descriptor.rs @@ -0,0 +1,122 @@ +use serde_json::json; + +use crate::tool::{ + RuntimeToolDescriptor, ToolApprovalCategory, ToolCapability, ToolDurability, + ToolExecutionCategory, ToolSideEffectLevel, + internal::{RuntimeDescriptorParts, build_runtime_descriptor}, +}; + +use super::RuntimeIntrinsicTool; + +pub(super) fn runtime_intrinsic_descriptor(tool: RuntimeIntrinsicTool) -> RuntimeToolDescriptor { + match tool { + RuntimeIntrinsicTool::Compact => build_runtime_descriptor(RuntimeDescriptorParts { + name: tool.to_string(), + description: "Compress older conversation context into a summary.".to_string(), + input_schema: json!({ + "type": "object", + "properties": {} + }), + capabilities: vec![ToolCapability::ContextCompaction], + side_effect_level: ToolSideEffectLevel::LocalState, + durability: ToolDurability::Persistent, + execution_category: ToolExecutionCategory::ExclusivePersistentMutation, + approval_category: ToolApprovalCategory::Default, + }), + RuntimeIntrinsicTool::Idle => build_runtime_descriptor(RuntimeDescriptorParts { + name: tool.to_string(), + description: "Yield the current turn and return to the teammate idle loop.".to_string(), + input_schema: json!({ + "type": "object", + "properties": {} + }), + capabilities: vec![ToolCapability::Delegation], + side_effect_level: ToolSideEffectLevel::LocalState, + durability: ToolDurability::Persistent, + execution_category: ToolExecutionCategory::Delegation, + approval_category: ToolApprovalCategory::Delegation, + }), + RuntimeIntrinsicTool::MemorySearch => build_runtime_descriptor(RuntimeDescriptorParts { + name: tool.to_string(), + description: "Search the current agent's long-term memory for additional recall." + .to_string(), + input_schema: json!({ + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "Memory query text" + }, + "limit": { + "type": "integer", + "description": "Maximum number of results to return" + } + }, + "required": ["query"] + }), + capabilities: vec![ToolCapability::ReadOnly], + side_effect_level: ToolSideEffectLevel::None, + durability: ToolDurability::ReplaySafe, + execution_category: ToolExecutionCategory::ReadOnlyParallel, + approval_category: ToolApprovalCategory::ReadOnly, + }), + RuntimeIntrinsicTool::MemoryPin => build_runtime_descriptor(RuntimeDescriptorParts { + name: tool.to_string(), + description: "Persist a fact in long-term memory for the current agent.".to_string(), + input_schema: json!({ + "type": "object", + "properties": { + "content": { + "type": "string", + "description": "Fact to remember" + } + }, + "required": ["content"] + }), + capabilities: vec![ToolCapability::Custom("memory_write".to_string())], + side_effect_level: ToolSideEffectLevel::LocalState, + durability: ToolDurability::Persistent, + execution_category: ToolExecutionCategory::ExclusivePersistentMutation, + approval_category: ToolApprovalCategory::Default, + }), + RuntimeIntrinsicTool::MemoryForget => build_runtime_descriptor(RuntimeDescriptorParts { + name: tool.to_string(), + description: "Forget a specific long-term memory record by id.".to_string(), + input_schema: json!({ + "type": "object", + "properties": { + "record_id": { + "type": "string", + "description": "Identifier of the memory record to forget" + } + }, + "required": ["record_id"] + }), + capabilities: vec![ToolCapability::Custom("memory_write".to_string())], + side_effect_level: ToolSideEffectLevel::LocalState, + durability: ToolDurability::Persistent, + execution_category: ToolExecutionCategory::ExclusivePersistentMutation, + approval_category: ToolApprovalCategory::Default, + }), + RuntimeIntrinsicTool::Task => build_runtime_descriptor(RuntimeDescriptorParts { + name: tool.to_string(), + description: "Spawn a fresh subagent to work a subtask and return a concise summary." + .to_string(), + input_schema: json!({ + "type": "object", + "properties": { + "prompt": { + "type": "string", + "description": "Delegated task prompt for the subagent" + } + }, + "required": ["prompt"] + }), + capabilities: vec![ToolCapability::Delegation], + side_effect_level: ToolSideEffectLevel::LocalState, + durability: ToolDurability::Ephemeral, + execution_category: ToolExecutionCategory::Delegation, + approval_category: ToolApprovalCategory::Delegation, + }), + } +} diff --git a/vendor/mentra/src/runtime/intrinsic/execute.rs b/vendor/mentra/src/runtime/intrinsic/execute.rs new file mode 100644 index 0000000..d7d0aad --- /dev/null +++ b/vendor/mentra/src/runtime/intrinsic/execute.rs @@ -0,0 +1,379 @@ +use serde_json::json; + +use crate::{ + ContentBlock, + agent::{Agent, AgentEvent, CompactionTrigger, SpawnedAgentStatus}, + memory::{MemorySearchMode, MemorySearchRequest}, + runtime::RunOptions, + tool::{ + ParallelToolContext, ToolCall, ToolContext, ToolResult, + internal::content_block_to_tool_result, + }, + transcript::{DelegationArtifact, DelegationEdge, DelegationKind, DelegationStatus}, +}; + +use super::{RuntimeIntrinsicTool, descriptor::runtime_intrinsic_descriptor}; + +pub(super) async fn execute_parallel( + tool: RuntimeIntrinsicTool, + ctx: ParallelToolContext, + input: serde_json::Value, +) -> ToolResult { + match tool { + RuntimeIntrinsicTool::MemorySearch => execute_memory_search(ctx, input).await, + _ => Err(format!( + "Tool '{}' does not support parallel execution", + runtime_intrinsic_descriptor(tool).provider.name + )), + } +} + +pub(super) async fn execute_mut( + tool: RuntimeIntrinsicTool, + ctx: ToolContext<'_>, + input: serde_json::Value, +) -> ToolResult { + match tool { + RuntimeIntrinsicTool::MemorySearch => execute_memory_search(ctx.into(), input).await, + _ => { + let call = ToolCall { + id: ctx.tool_call_id.clone(), + name: runtime_intrinsic_descriptor(tool).provider.name, + input, + }; + let child_options = ctx.child_run_options(); + let block = match tool { + RuntimeIntrinsicTool::Compact => execute_compact(ctx.agent, call).await, + RuntimeIntrinsicTool::Idle => execute_idle(ctx.agent, call), + RuntimeIntrinsicTool::MemorySearch => unreachable!("handled above"), + RuntimeIntrinsicTool::MemoryPin => execute_memory_pin(ctx, call), + RuntimeIntrinsicTool::MemoryForget => execute_memory_forget(ctx, call), + RuntimeIntrinsicTool::Task => execute_task(ctx.agent, call, child_options).await, + }; + content_block_to_tool_result("Runtime intrinsic", block) + } + } +} + +fn execute_idle(agent: &mut Agent, call: ToolCall) -> ContentBlock { + agent.request_idle(); + ContentBlock::ToolResult { + tool_use_id: call.id, + content: "Yielding to the teammate idle loop.".into(), + is_error: false, + } +} + +fn execute_memory_pin(ctx: ToolContext<'_>, call: ToolCall) -> ContentBlock { + if !ctx.agent.config().memory.write_tools_enabled { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: "Memory write tools are disabled for this agent.".into(), + is_error: true, + }; + } + + let Some(content) = call + .input + .get("content") + .and_then(|value| value.as_str()) + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: "Invalid memory_pin input: content is required.".into(), + is_error: true, + }; + }; + + match ctx + .agent + .memory_engine() + .pin(ctx.agent.id(), ctx.agent.memory_revision(), content) + { + Ok(record) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Pinned memory {}.", record.record_id).into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to pin memory: {error}").into(), + is_error: true, + }, + } +} + +fn execute_memory_forget(ctx: ToolContext<'_>, call: ToolCall) -> ContentBlock { + if !ctx.agent.config().memory.write_tools_enabled { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: "Memory write tools are disabled for this agent.".into(), + is_error: true, + }; + } + + let Some(record_id) = call + .input + .get("record_id") + .and_then(|value| value.as_str()) + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: "Invalid memory_forget input: record_id is required.".into(), + is_error: true, + }; + }; + + match ctx.agent.memory_engine().forget(ctx.agent.id(), record_id) { + Ok(true) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Forgot memory {record_id}.").into(), + is_error: false, + }, + Ok(false) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Memory record {record_id} was not found for this agent.").into(), + is_error: true, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to forget memory: {error}").into(), + is_error: true, + }, + } +} + +async fn execute_memory_search(ctx: ParallelToolContext, input: serde_json::Value) -> ToolResult { + let Some(query) = input + .get("query") + .and_then(|value| value.as_str()) + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return Err("Invalid memory_search input: query is required.".to_string()); + }; + + let configured_limit = ctx + .runtime + .agent_config(&ctx.agent_id)? + .memory_tool_search_limit; + let requested_limit = input + .get("limit") + .and_then(|value| value.as_u64()) + .unwrap_or(configured_limit as u64) as usize; + let limit = requested_limit.min(configured_limit).min(10); + + match ctx + .runtime + .memory_engine() + .search(MemorySearchRequest { + agent_id: ctx.agent_id.clone(), + query: query.to_string(), + limit, + char_budget: None, + mode: MemorySearchMode::Tool, + filter: crate::memory::MemoryListFilter::default(), + }) + .await + { + Ok(hits) => { + let results = hits + .into_iter() + .map(|hit| { + json!({ + "id": hit.record_id, + "kind": hit.kind, + "content": hit.content, + "score": hit.score, + "timestamp": hit.created_at, + "source": hit.source, + "why_retrieved": hit.why_retrieved, + }) + }) + .collect::>(); + Ok(serde_json::to_string_pretty(&results).unwrap_or_else(|_| "[]".to_string())) + } + Err(error) => Err(format!("Memory search failed: {error}")), + } +} + +async fn execute_compact(agent: &mut Agent, call: ToolCall) -> ContentBlock { + match agent + .compact_history( + agent.history().len().saturating_sub(1), + CompactionTrigger::Manual, + ) + .await + { + Ok(Some(details)) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!( + "Context compacted. Transcript saved to {}", + details.transcript_path.display() + ) + .into(), + is_error: false, + }, + Ok(None) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: "Context compaction skipped because there was no older history to summarize." + .into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Context compaction failed: {error}").into(), + is_error: true, + }, + } +} + +/// Runs a delegated task on a disposable subagent. +/// +/// `options` are the parent run's [`RunOptions::child`]: the delegated run +/// reports its token usage into the parent's accounting handle and shares the +/// parent's cancellation, stop, and deadline. Driving the child on +/// `RunOptions::default()` instead would let delegated spend escape the +/// parent's `token_budget` and leave the child running after a parent cancel. +async fn execute_task(agent: &mut Agent, call: ToolCall, options: RunOptions) -> ContentBlock { + match crate::agent::parse_task_input(call.input) { + Ok(prompt) => { + let task_summary = prompt.clone(); + let mut child = match agent.spawn_subagent() { + Ok(child) => child, + Err(error) => { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to spawn subagent: {error}").into(), + is_error: true, + }; + } + }; + let child_id = child.id().to_string(); + let child_name = child.name().to_string(); + let child_model = child.model().to_string(); + let edge = Some(DelegationEdge { + kind: DelegationKind::Subagent, + local_agent_id: agent.id().to_string(), + remote_agent_id: child_id.clone(), + }); + let _ = agent.record_delegation_request( + format!( + "\n{prompt}\n" + ), + DelegationArtifact { + kind: DelegationKind::Subagent, + agent_id: child_id.clone(), + agent_name: child_name.clone(), + role: Some("subagent".to_string()), + status: DelegationStatus::Requested, + task_summary: task_summary.clone(), + result_summary: None, + artifacts: Vec::new(), + }, + edge.clone(), + ); + agent.sync_memory_snapshot(); + let started = agent.register_subagent(&child); + agent.emit_event(AgentEvent::SubagentSpawned { agent: started }); + + // A subagent has its own event bus, so a parent's observer would + // otherwise see none of the delegated spend that now counts against + // the parent's `token_budget`. Relaying just `UsageReport` keeps the + // parent's stream summing to the same total the shared accounting + // handle reports. The guard must outlive the child's run. + let parent_events = agent.event_sender(); + let _usage_relay = child.register_event_tap(move |event| { + if matches!(event, AgentEvent::UsageReport { .. }) { + parent_events.send(event.clone()); + } + }); + + match Box::pin(child.run(vec![ContentBlock::Text { text: prompt }], options)).await { + Ok(message) => { + let result_summary = if message.text().is_empty() { + child.final_text_summary() + } else { + message.text() + }; + let _ = agent.record_delegation_result( + format!( + "\n{result_summary}\n" + ), + DelegationArtifact { + kind: DelegationKind::Subagent, + agent_id: child_id.clone(), + agent_name: child_name.clone(), + role: Some("subagent".to_string()), + status: DelegationStatus::Finished, + task_summary: task_summary.clone(), + result_summary: Some(result_summary.clone()), + artifacts: Vec::new(), + }, + edge.clone(), + ); + agent.sync_memory_snapshot(); + if let Some(finished) = + agent.finish_subagent(child.id(), SpawnedAgentStatus::Finished) + { + agent.emit_event(AgentEvent::SubagentFinished { agent: finished }); + } + if let Err(error) = agent.refresh_tasks_from_disk() { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Task refresh failed: {error}").into(), + is_error: true, + }; + } + + ContentBlock::ToolResult { + tool_use_id: call.id, + content: result_summary.into(), + is_error: false, + } + } + Err(error) => { + let error_text = error.to_string(); + let _ = agent.record_delegation_result( + format!( + "\n{error_text}\n" + ), + DelegationArtifact { + kind: DelegationKind::Subagent, + agent_id: child_id, + agent_name: child_name, + role: Some("subagent".to_string()), + status: DelegationStatus::Failed, + task_summary, + result_summary: Some(error_text.clone()), + artifacts: Vec::new(), + }, + edge, + ); + agent.sync_memory_snapshot(); + if let Some(finished) = agent + .finish_subagent(child.id(), SpawnedAgentStatus::Failed(error_text.clone())) + { + agent.emit_event(AgentEvent::SubagentFinished { agent: finished }); + } + let _ = agent.refresh_tasks_from_disk(); + + ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Subagent failed: {error_text}").into(), + is_error: true, + } + } + } + } + Err(content) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: content.into(), + is_error: true, + }, + } +} diff --git a/vendor/mentra/src/runtime/skill.rs b/vendor/mentra/src/runtime/skill.rs new file mode 100644 index 0000000..02f9ec5 --- /dev/null +++ b/vendor/mentra/src/runtime/skill.rs @@ -0,0 +1,415 @@ +use std::{ + collections::BTreeMap, + fs, + path::{Path, PathBuf}, +}; + +use serde::Deserialize; +use thiserror::Error; + +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub(crate) struct SkillLoader { + skills: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct SkillEntry { + description: String, + body: String, + path: PathBuf, +} + +/// A loaded skill, without its body. +/// +/// Name and description are what a host needs to show a skill set to a person +/// — in a client UI, as protocol commands, in a run's log, or in a test +/// asserting the expected skills loaded. The body stays behind `load_skill`, +/// which is what keeps skills cheap in context: descriptions are always +/// present, bodies arrive only when asked for. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SkillInfo { + pub name: String, + pub description: String, + /// The `SKILL.md` this came from. With several roots registered, this is + /// how a host tells which one won. + pub path: PathBuf, +} + +#[derive(Debug, Clone, PartialEq, Eq, Error)] +pub enum SkillLoadError { + #[error("failed to read skills directory {path}: {message}")] + ReadDir { path: PathBuf, message: String }, + #[error("failed to read skill file {path}: {message}")] + ReadFile { path: PathBuf, message: String }, + #[error("invalid skill frontmatter in {path}: {message}")] + InvalidFrontmatter { path: PathBuf, message: String }, + #[error("duplicate skill name '{name}' in {first_path} and {second_path}")] + DuplicateSkillName { + name: String, + first_path: PathBuf, + second_path: PathBuf, + }, +} + +#[derive(Debug, Clone, Default, Deserialize)] +struct SkillFrontmatter { + name: Option, + description: Option, +} + +impl SkillLoader { + pub(crate) fn from_dir(path: impl AsRef) -> Result { + let root = path.as_ref().to_path_buf(); + let mut files = Vec::new(); + collect_skill_files(&root, &mut files)?; + files.sort(); + + let mut skills = BTreeMap::new(); + let mut skill_paths = BTreeMap::new(); + + for file in files { + let raw = fs::read_to_string(&file).map_err(|error| SkillLoadError::ReadFile { + path: file.clone(), + message: error.to_string(), + })?; + let (meta, body) = parse_skill_file(&file, &raw)?; + + let fallback_name = file + .parent() + .and_then(Path::file_name) + .and_then(|value| value.to_str()) + .unwrap_or("skill"); + let name = meta + .name + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or(fallback_name) + .to_string(); + + if let Some(first_path) = skill_paths.insert(name.clone(), file.clone()) { + return Err(SkillLoadError::DuplicateSkillName { + name, + first_path, + second_path: file, + }); + } + + let description = meta.description.unwrap_or_default().trim().to_string(); + skills.insert( + name, + SkillEntry { + description, + body, + path: file, + }, + ); + } + + Ok(Self { skills }) + } + + /// Folds in skills from a lower-precedence root. + /// + /// A name already defined here wins, so roots registered earlier shadow + /// later ones — the same rule `PATH` uses, and the one that lets a project + /// override a personal skill by name. Within a single root a repeated name + /// is still [`SkillLoadError::DuplicateSkillName`], because there it is a + /// mistake rather than an intent. + pub(crate) fn merge_weaker(&mut self, weaker: SkillLoader) { + for (name, entry) in weaker.skills { + self.skills.entry(name).or_insert(entry); + } + } + + /// Every loaded skill, name-ordered, without bodies. + pub(crate) fn infos(&self) -> Vec { + self.skills + .iter() + .map(|(name, entry)| SkillInfo { + name: name.clone(), + description: entry.description.clone(), + path: entry.path.clone(), + }) + .collect() + } + + pub(crate) fn get_descriptions(&self) -> String { + if self.skills.is_empty() { + return String::new(); + } + + let mut lines = vec!["Skills available:".to_string()]; + for (name, skill) in &self.skills { + lines.push(format!(" - {name}: {}", skill.description)); + } + lines.push( + "Use the load_skill tool only when one of these skills is relevant to the task." + .to_string(), + ); + lines.join("\n") + } + + pub(crate) fn get_content(&self, name: &str) -> Result { + let Some(skill) = self.skills.get(name) else { + return Err(format!("Unknown skill '{name}'")); + }; + + let body = skill.body.trim_end_matches(['\n', '\r']); + Ok(format!("\n{body}\n")) + } +} + +fn collect_skill_files(path: &Path, files: &mut Vec) -> Result<(), SkillLoadError> { + let entries = fs::read_dir(path).map_err(|error| SkillLoadError::ReadDir { + path: path.to_path_buf(), + message: error.to_string(), + })?; + + for entry in entries { + let entry = entry.map_err(|error| SkillLoadError::ReadDir { + path: path.to_path_buf(), + message: error.to_string(), + })?; + let entry_path = entry.path(); + let file_type = entry.file_type().map_err(|error| SkillLoadError::ReadDir { + path: entry_path.clone(), + message: error.to_string(), + })?; + + if file_type.is_dir() { + collect_skill_files(&entry_path, files)?; + } else if file_type.is_file() && entry.file_name() == "SKILL.md" { + files.push(entry_path); + } + } + + Ok(()) +} + +fn parse_skill_file(path: &Path, raw: &str) -> Result<(SkillFrontmatter, String), SkillLoadError> { + let Some(opening_len) = raw + .strip_prefix("---\r\n") + .map(|_| 5) + .or_else(|| raw.strip_prefix("---\n").map(|_| 4)) + else { + return Ok((SkillFrontmatter::default(), raw.to_string())); + }; + + let rest = &raw[opening_len..]; + let mut cursor = 0usize; + for segment in rest.split_inclusive('\n') { + let line = segment.trim_end_matches(['\n', '\r']); + if line == "---" { + let frontmatter = &rest[..cursor]; + let body = &rest[cursor + segment.len()..]; + let meta = serde_yaml_ng::from_str(frontmatter).map_err(|error| { + SkillLoadError::InvalidFrontmatter { + path: path.to_path_buf(), + message: error.to_string(), + } + })?; + return Ok((meta, body.to_string())); + } + cursor += segment.len(); + } + + if rest[cursor..].trim_end_matches('\r') == "---" { + let frontmatter = &rest[..cursor]; + let meta = serde_yaml_ng::from_str(frontmatter).map_err(|error| { + SkillLoadError::InvalidFrontmatter { + path: path.to_path_buf(), + message: error.to_string(), + } + })?; + return Ok((meta, String::new())); + } + + Err(SkillLoadError::InvalidFrontmatter { + path: path.to_path_buf(), + message: "missing closing frontmatter delimiter".to_string(), + }) +} + +#[cfg(test)] +mod tests { + use std::{ + fs, + path::{Path, PathBuf}, + sync::atomic::{AtomicU64, Ordering}, + time::{SystemTime, UNIX_EPOCH}, + }; + + use super::{SkillLoadError, SkillLoader}; + + static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + + #[test] + fn parses_frontmatter_and_strips_it_from_content() { + let root = temp_skills_dir("frontmatter"); + write_skill( + &root, + "git", + "---\nname: git\ndescription: Git helpers\n---\nStep 1\nStep 2\n", + ); + + let loader = SkillLoader::from_dir(&root).expect("load skills"); + + assert_eq!( + loader.get_descriptions(), + "Skills available:\n - git: Git helpers\nUse the load_skill tool only when one of these skills is relevant to the task." + ); + assert_eq!( + loader.get_content("git").expect("git skill"), + "\nStep 1\nStep 2\n" + ); + } + + #[test] + fn falls_back_to_directory_name_when_name_is_missing() { + let root = temp_skills_dir("fallback-name"); + write_skill( + &root, + "pdf", + "---\ndescription: Process PDFs\n---\nRead pages\n", + ); + + let loader = SkillLoader::from_dir(&root).expect("load skills"); + + assert!(loader.get_descriptions().contains(" - pdf: Process PDFs")); + assert!(loader.get_content("pdf").is_ok()); + } + + #[test] + fn renders_descriptions_in_sorted_order() { + let root = temp_skills_dir("sorted"); + write_skill( + &root, + "b-skill", + "---\nname: zebra\ndescription: Last\n---\nB\n", + ); + write_skill( + &root, + "a-skill", + "---\nname: alpha\ndescription: First\n---\nA\n", + ); + + let loader = SkillLoader::from_dir(&root).expect("load skills"); + + assert_eq!( + loader.get_descriptions(), + "Skills available:\n - alpha: First\n - zebra: Last\nUse the load_skill tool only when one of these skills is relevant to the task." + ); + } + + #[test] + fn rejects_duplicate_skill_names() { + let root = temp_skills_dir("duplicate"); + write_skill(&root, "one", "---\nname: shared\n---\nA\n"); + write_skill(&root, "two", "---\nname: shared\n---\nB\n"); + + let error = SkillLoader::from_dir(&root).expect_err("duplicate error"); + + assert!(matches!( + error, + SkillLoadError::DuplicateSkillName { ref name, .. } if name == "shared" + )); + } + + #[test] + fn rejects_malformed_frontmatter() { + let root = temp_skills_dir("invalid-frontmatter"); + write_skill(&root, "broken", "---\nname: [oops\n---\nBody\n"); + + let error = SkillLoader::from_dir(&root).expect_err("frontmatter error"); + + assert!(matches!(error, SkillLoadError::InvalidFrontmatter { .. })); + assert!(error.to_string().contains("invalid skill frontmatter")); + } + + #[test] + fn a_weaker_root_only_fills_in_names_the_stronger_one_lacks() { + let strong = temp_skills_dir("merge-strong"); + write_skill(&strong, "review", "---\nname: review\n---\nProject rules\n"); + let weak = temp_skills_dir("merge-weak"); + write_skill(&weak, "review", "---\nname: review\n---\nPersonal rules\n"); + write_skill(&weak, "deploy", "---\nname: deploy\n---\nPersonal deploy\n"); + + let mut loader = SkillLoader::from_dir(&strong).expect("strong root loads"); + loader.merge_weaker(SkillLoader::from_dir(&weak).expect("weak root loads")); + + assert!( + loader + .get_content("review") + .expect("review present") + .contains("Project rules"), + "the stronger root must win a name collision" + ); + assert!( + loader.get_content("deploy").is_ok(), + "a name only the weaker root defines must still load" + ); + } + + #[test] + fn merging_reports_which_file_each_skill_came_from() { + let strong = temp_skills_dir("merge-infos-strong"); + write_skill( + &strong, + "review", + "---\nname: review\ndescription: D1\n---\nA\n", + ); + let weak = temp_skills_dir("merge-infos-weak"); + write_skill( + &weak, + "deploy", + "---\nname: deploy\ndescription: D2\n---\nB\n", + ); + + let mut loader = SkillLoader::from_dir(&strong).expect("loads"); + loader.merge_weaker(SkillLoader::from_dir(&weak).expect("loads")); + let infos = loader.infos(); + + assert_eq!(infos.len(), 2); + // Name-ordered, so `deploy` precedes `review`. + assert_eq!(infos[0].name, "deploy"); + assert_eq!(infos[0].description, "D2"); + assert!(infos[0].path.starts_with(&weak)); + assert_eq!(infos[1].name, "review"); + assert_eq!(infos[1].description, "D1"); + assert!(infos[1].path.starts_with(&strong)); + } + + #[test] + fn infos_omit_bodies() { + let root = temp_skills_dir("infos-no-body"); + write_skill( + &root, + "one", + "---\nname: one\ndescription: short\n---\nSECRET BODY\n", + ); + + let infos = SkillLoader::from_dir(&root).expect("loads").infos(); + + let rendered = format!("{infos:?}"); + assert!(!rendered.contains("SECRET BODY")); + } + + fn temp_skills_dir(label: &str) -> PathBuf { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = + std::env::temp_dir().join(format!("mentra-skill-tests-{label}-{timestamp}-{unique}")); + fs::create_dir_all(&path).expect("create temp dir"); + path + } + + fn write_skill(root: &Path, name: &str, content: &str) { + let skill_dir = root.join(name); + fs::create_dir_all(&skill_dir).expect("create skill dir"); + fs::write(skill_dir.join("SKILL.md"), content).expect("write skill"); + } +} diff --git a/vendor/mentra/src/runtime/store.rs b/vendor/mentra/src/runtime/store.rs new file mode 100644 index 0000000..612f3ef --- /dev/null +++ b/vendor/mentra/src/runtime/store.rs @@ -0,0 +1,2432 @@ +use std::{ + collections::HashSet, + path::{Path, PathBuf}, + sync::atomic::{AtomicU64, Ordering}, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use rusqlite::{Connection, OptionalExtension, TransactionBehavior, params}; +use serde::{Deserialize, Serialize, de::DeserializeOwned}; + +use crate::{ + agent::{AgentConfig, AgentStatus, SpawnedAgentSummary, TeammateIdentity}, + background::{ + BackgroundNotification, BackgroundStore, BackgroundTaskStatus, BackgroundTaskSummary, + }, + memory::journal::AgentMemoryState, + memory::{ + MemoryCursor, MemoryListCursor, MemoryListPage, MemoryListRequest, MemoryListSort, + MemoryRecord, MemorySearchRequest, MemoryStore, + }, + provider::ProviderId, + runtime::TaskItem, + session::PermissionRuleScope, + session::permission::{RememberedRule, RuleKey}, + team::{TeamMemberSummary, TeamMessage, TeamProtocolRequestSummary, TeamStore}, +}; + +use super::error::RuntimeError; + +static NEXT_STORE_ID: AtomicU64 = AtomicU64::new(1); +#[cfg(test)] +static NEXT_TEST_STORE_ID: AtomicU64 = AtomicU64::new(1); + +const DELIVERY_PENDING: i64 = 0; +const DELIVERY_INFLIGHT: i64 = 1; +const DELIVERY_ACKED: i64 = 2; + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PersistedAgentRecord { + pub(crate) id: String, + pub(crate) runtime_identifier: String, + pub(crate) name: String, + pub(crate) model: String, + pub(crate) provider_id: ProviderId, + pub(crate) config: AgentConfig, + pub(crate) hidden_tools: HashSet, + pub(crate) max_rounds: Option, + pub(crate) teammate_identity: Option, + pub(crate) rounds_since_task: usize, + pub(crate) idle_requested: bool, + pub(crate) status: AgentStatus, + pub(crate) subagents: Vec, +} + +#[derive(Debug, Clone)] +pub struct LoadedAgentState { + pub(crate) record: PersistedAgentRecord, + pub(crate) memory: AgentMemoryState, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TaskStateSnapshot { + pub(crate) tasks: Vec, +} + +/// Persistence backend for agent records and working-memory snapshots. +/// +/// Custom runtime backends implement this trait to store durable agent identity, +/// configuration, and transcript state. +pub trait AgentStore: Send + Sync { + /// Returns whether runtime-managed auxiliary artifacts may be written to + /// disk for agents backed by this store. + /// + /// Persistent stores allow artifacts by default. Volatile stores override + /// this capability so features such as full tool-output spilling preserve + /// their no-durable-trace contract. + fn allows_disk_artifacts(&self) -> bool { + true + } + + fn prepare_recovery(&self) -> Result<(), RuntimeError>; + fn create_agent( + &self, + record: &PersistedAgentRecord, + memory: &AgentMemoryState, + ) -> Result<(), RuntimeError>; + fn save_agent_record(&self, record: &PersistedAgentRecord) -> Result<(), RuntimeError>; + fn save_agent_memory( + &self, + agent_id: &str, + memory: &AgentMemoryState, + ) -> Result<(), RuntimeError>; + fn load_agent(&self, agent_id: &str) -> Result, RuntimeError>; + fn list_agents(&self) -> Result, RuntimeError>; + fn list_agents_by_runtime( + &self, + runtime_identifier: &str, + ) -> Result, RuntimeError>; +} + +/// Persistence backend for tracked agent runs. +/// +/// This trait stores lifecycle transitions for turns and interrupted runs. +pub trait RunStore: Send + Sync { + fn start_run(&self, agent_id: &str) -> Result; + fn update_run_state( + &self, + run_id: &str, + state: &str, + error: Option<&str>, + ) -> Result<(), RuntimeError>; + fn finish_run(&self, run_id: &str) -> Result<(), RuntimeError>; + fn fail_run(&self, run_id: &str, error: &str) -> Result<(), RuntimeError>; +} + +/// Persistence backend for the dependency-aware task board. +/// +/// Task persistence is intentionally separate so applications can replace the +/// task board without reimplementing unrelated runtime storage. +pub trait TaskStore: Send + Sync { + fn load_tasks(&self, namespace: &Path) -> Result, RuntimeError>; + fn capture_tasks(&self, namespace: &Path) -> Result; + fn restore_tasks( + &self, + namespace: &Path, + snapshot: &TaskStateSnapshot, + ) -> Result<(), RuntimeError>; + fn replace_tasks(&self, namespace: &Path, tasks: &[TaskItem]) -> Result<(), RuntimeError>; + + /// Applies one read-modify-write operation to a namespace. + /// + /// The callback form keeps this method object-safe, so runtime code can use + /// it through `dyn TaskStore`. The default preserves source compatibility + /// for external stores by composing [`TaskStore::load_tasks`] and + /// [`TaskStore::replace_tasks`], but that fallback cannot promise + /// serialization across concurrent writers. Stores that can provide a + /// transaction or lock should override this method. + /// + /// If `mutation` returns an error, the modified task vector must not be + /// installed by overrides. + fn mutate( + &self, + namespace: &Path, + mutation: &mut dyn FnMut(&mut Vec) -> Result<(), RuntimeError>, + ) -> Result<(), RuntimeError> { + let mut tasks = self.load_tasks(namespace)?; + mutation(&mut tasks)?; + self.replace_tasks(namespace, &tasks) + } +} + +/// Persistence backend for runtime audit hooks. +pub trait AuditStore: Send + Sync { + fn record_audit_event( + &self, + scope: &str, + event_type: &str, + payload: serde_json::Value, + ) -> Result<(), RuntimeError>; +} + +/// Persistence backend for runtime leases. +/// +/// Leases coordinate exclusive ownership when multiple runtime processes may try +/// to resume the same persisted agents. +pub trait LeaseStore: Send + Sync { + fn acquire_lease(&self, key: &str, owner: &str, ttl: Duration) -> Result; + fn release_lease(&self, key: &str, owner: &str) -> Result<(), RuntimeError>; +} + +/// Persistence backend for remembered permission rules. +/// +/// Permission rules survive session restarts when backed by a persistent store. +/// +/// The `project_id` parameter is an opaque string supplied by the consumer and +/// used to associate rules with a project for cross-session inheritance. +/// Mentra does not interpret its value. +pub trait PermissionRuleStore: Send + Sync { + /// Persists the provided permission rules for a session, replacing any + /// existing session-scoped rules. `project_id` is stored alongside each + /// rule so that project-scoped rules can later be retrieved by other + /// sessions that share the same project. + fn save_rules( + &self, + session_id: &str, + project_id: Option<&str>, + rules: &[RememberedRule], + ) -> Result<(), RuntimeError>; + + /// Loads all persisted permission rules that apply to the given session. + /// + /// The returned set is the union of: + /// - Session-scoped rules where `session_id` matches. + /// - Project-scoped rules where `project_id` matches (when provided). + /// - Global-scoped rules (always included). + fn load_rules( + &self, + session_id: &str, + project_id: Option<&str>, + ) -> Result, RuntimeError>; + + /// Removes all persisted permission rules for a session. + fn clear_rules(&self, session_id: &str) -> Result<(), RuntimeError>; +} + +/// Full persistence backend used by the runtime. +/// +/// `RuntimeStore` is a composition trait over the narrower persistence seams +/// plus the collaboration and memory stores. Custom backends can implement the +/// smaller traits directly and then satisfy `RuntimeStore` automatically. +pub trait RuntimeStore: + AgentStore + + RunStore + + TaskStore + + AuditStore + + LeaseStore + + PermissionRuleStore + + TeamStore + + BackgroundStore + + MemoryStore + + Send + + Sync +{ +} + +impl RuntimeStore for T where + T: AgentStore + + RunStore + + TaskStore + + AuditStore + + LeaseStore + + PermissionRuleStore + + TeamStore + + BackgroundStore + + MemoryStore + + Send + + Sync +{ +} + +impl TeamStore for SqliteRuntimeStore { + fn unread_team_count(&self, team_dir: &Path, agent_name: &str) -> Result { + let conn = self.open()?; + let count = conn + .query_row( + "SELECT COUNT(*) FROM team_inbox WHERE team_dir = ?1 AND recipient = ?2 AND delivery_state = ?3", + params![Self::team_key(team_dir), agent_name, DELIVERY_PENDING], + |row| row.get::<_, i64>(0), + ) + .map_err(sqlite_error)?; + Ok(count as usize) + } + + fn load_team_members(&self, team_dir: &Path) -> Result, RuntimeError> { + let conn = self.open()?; + let mut stmt = conn + .prepare("SELECT summary_json FROM team_members WHERE team_dir = ?1 ORDER BY name") + .map_err(sqlite_error)?; + let rows = stmt + .query_map(params![Self::team_key(team_dir)], |row| { + row.get::<_, String>(0) + }) + .map_err(sqlite_error)?; + let mut members = Vec::new(); + for row in rows { + members.push(from_json(&row.map_err(sqlite_error)?)?); + } + Ok(members) + } + + fn upsert_team_member( + &self, + team_dir: &Path, + summary: &TeamMemberSummary, + ) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + r#" + INSERT INTO team_members (team_dir, name, summary_json) + VALUES (?1, ?2, ?3) + ON CONFLICT(team_dir, name) DO UPDATE SET summary_json = excluded.summary_json + "#, + params![Self::team_key(team_dir), summary.name, to_json(summary)?], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn read_team_inbox( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result, RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + let team_key = Self::team_key(team_dir); + let ids_and_payloads = { + let mut stmt = tx + .prepare( + "SELECT id, payload_json FROM team_inbox WHERE team_dir = ?1 AND recipient = ?2 AND delivery_state = ?3 ORDER BY created_at, id", + ) + .map_err(sqlite_error)?; + stmt.query_map(params![team_key, agent_name, DELIVERY_PENDING], |row| { + Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)) + }) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error)? + }; + + for (id, _) in &ids_and_payloads { + tx.execute( + "UPDATE team_inbox SET delivery_state = ?2 WHERE id = ?1", + params![id, DELIVERY_INFLIGHT], + ) + .map_err(sqlite_error)?; + } + tx.commit().map_err(sqlite_error)?; + + ids_and_payloads + .into_iter() + .map(|(_, payload)| from_json(&payload)) + .collect() + } + + fn ack_team_inbox(&self, team_dir: &Path, agent_name: &str) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "UPDATE team_inbox SET delivery_state = ?3 WHERE team_dir = ?1 AND recipient = ?2 AND delivery_state = ?4", + params![Self::team_key(team_dir), agent_name, DELIVERY_ACKED, DELIVERY_INFLIGHT], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn requeue_team_inbox(&self, team_dir: &Path, agent_name: &str) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "UPDATE team_inbox SET delivery_state = ?3 WHERE team_dir = ?1 AND recipient = ?2 AND delivery_state = ?4", + params![Self::team_key(team_dir), agent_name, DELIVERY_PENDING, DELIVERY_INFLIGHT], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn append_team_message( + &self, + team_dir: &Path, + recipient: &str, + message: &TeamMessage, + ) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "INSERT INTO team_inbox (id, team_dir, recipient, payload_json, delivery_state, created_at) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + next_id("teammsg"), + Self::team_key(team_dir), + recipient, + to_json(message)?, + DELIVERY_PENDING, + now_secs(), + ], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn load_team_requests( + &self, + team_dir: &Path, + ) -> Result, RuntimeError> { + let conn = self.open()?; + let mut stmt = conn + .prepare( + "SELECT payload_json FROM team_requests WHERE team_dir = ?1 ORDER BY created_at, request_id", + ) + .map_err(sqlite_error)?; + let rows = stmt + .query_map(params![Self::team_key(team_dir)], |row| { + row.get::<_, String>(0) + }) + .map_err(sqlite_error)?; + let mut requests = Vec::new(); + for row in rows { + requests.push(from_json(&row.map_err(sqlite_error)?)?); + } + Ok(requests) + } + + fn upsert_team_request( + &self, + team_dir: &Path, + request: &TeamProtocolRequestSummary, + ) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + r#" + INSERT INTO team_requests (request_id, team_dir, payload_json, created_at) + VALUES (?1, ?2, ?3, ?4) + ON CONFLICT(request_id) DO UPDATE SET + team_dir = excluded.team_dir, + payload_json = excluded.payload_json + "#, + params![ + request.request_id, + Self::team_key(team_dir), + to_json(request)?, + request.created_at as i64, + ], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn list_team_agent_names(&self, team_dir: &Path) -> Result, RuntimeError> { + let conn = self.open()?; + let mut stmt = conn + .prepare("SELECT name FROM agents WHERE team_dir = ?1 ORDER BY name") + .map_err(sqlite_error)?; + stmt.query_map(params![Self::team_key(team_dir)], |row| { + row.get::<_, String>(0) + }) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error) + } +} + +impl BackgroundStore for SqliteRuntimeStore { + fn load_background_tasks( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + let conn = self.open()?; + let mut stmt = conn + .prepare( + "SELECT payload_json FROM background_jobs WHERE agent_id = ?1 ORDER BY created_at, id", + ) + .map_err(sqlite_error)?; + let rows = stmt + .query_map(params![agent_id], |row| row.get::<_, String>(0)) + .map_err(sqlite_error)?; + let mut tasks = Vec::new(); + for row in rows { + tasks.push(from_json(&row.map_err(sqlite_error)?)?); + } + Ok(tasks) + } + + fn upsert_background_task( + &self, + agent_id: &str, + task: &BackgroundTaskSummary, + notification_state: i64, + ) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + r#" + INSERT INTO background_jobs (agent_id, id, payload_json, notification_state, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?5) + ON CONFLICT(agent_id, id) DO UPDATE SET + payload_json = excluded.payload_json, + notification_state = excluded.notification_state, + updated_at = excluded.updated_at + "#, + params![agent_id, task.id, to_json(task)?, notification_state, now_secs()], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn drain_background_notifications( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + let jobs = { + let mut stmt = tx + .prepare( + "SELECT id, payload_json FROM background_jobs WHERE agent_id = ?1 AND notification_state = ?2 ORDER BY updated_at, id", + ) + .map_err(sqlite_error)?; + stmt.query_map(params![agent_id, DELIVERY_PENDING], |row| { + Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)) + }) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error)? + }; + for (id, _) in &jobs { + tx.execute( + "UPDATE background_jobs SET notification_state = ?3 WHERE agent_id = ?1 AND id = ?2", + params![agent_id, id, DELIVERY_INFLIGHT], + ) + .map_err(sqlite_error)?; + } + tx.commit().map_err(sqlite_error)?; + + jobs.into_iter() + .map(|(_, payload)| { + let task: BackgroundTaskSummary = from_json(&payload)?; + Ok(BackgroundNotification { + task_id: task.id, + command: task.command, + cwd: task.cwd, + status: task.status, + output_preview: task + .output_preview + .unwrap_or_else(|| "(no output)".to_string()), + }) + }) + .collect() + } + + fn has_pending_background_notifications(&self, agent_id: &str) -> Result { + let conn = self.open()?; + let exists = conn + .query_row( + "SELECT EXISTS(SELECT 1 FROM background_jobs WHERE agent_id = ?1 AND notification_state IN (?2, ?3))", + params![agent_id, DELIVERY_PENDING, DELIVERY_INFLIGHT], + |row| row.get::<_, i64>(0), + ) + .map_err(sqlite_error)?; + Ok(exists != 0) + } + + fn has_deliverable_background_notifications( + &self, + agent_id: &str, + ) -> Result { + let conn = self.open()?; + let exists = conn + .query_row( + "SELECT EXISTS(SELECT 1 FROM background_jobs WHERE agent_id = ?1 AND notification_state = ?2)", + params![agent_id, DELIVERY_PENDING], + |row| row.get::<_, i64>(0), + ) + .map_err(sqlite_error)?; + Ok(exists != 0) + } + + fn ack_background_notifications(&self, agent_id: &str) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "UPDATE background_jobs SET notification_state = ?2 WHERE agent_id = ?1 AND notification_state = ?3", + params![agent_id, DELIVERY_ACKED, DELIVERY_INFLIGHT], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn requeue_background_notifications(&self, agent_id: &str) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "UPDATE background_jobs SET notification_state = ?2 WHERE agent_id = ?1 AND notification_state = ?3", + params![agent_id, DELIVERY_PENDING, DELIVERY_INFLIGHT], + ) + .map_err(sqlite_error)?; + Ok(()) + } +} + +#[derive(Clone)] +/// SQLite-backed [`RuntimeStore`] implementation used by default. +pub struct SqliteRuntimeStore { + path: PathBuf, +} + +impl Default for SqliteRuntimeStore { + fn default() -> Self { + Self::new(Self::default_path()) + } +} + +impl SqliteRuntimeStore { + /// Returns the default SQLite path used when no explicit store path is provided. + pub fn default_path() -> PathBuf { + default_store_dir().join("runtime.sqlite") + } + + /// Returns the default directory used for Mentra runtime stores. + pub fn default_directory() -> PathBuf { + default_store_dir() + } + + /// Creates a SQLite runtime store in the default directory using a runtime-scoped filename. + pub fn for_runtime_identifier(runtime_identifier: &str) -> Self { + Self::new(Self::path_for_runtime_identifier(runtime_identifier)) + } + + /// Returns the default SQLite path for a specific runtime identifier. + pub fn path_for_runtime_identifier(runtime_identifier: &str) -> PathBuf { + Self::default_directory().join(format!( + "runtime-{}.sqlite", + encode_runtime_identifier(runtime_identifier) + )) + } + + /// Lists runtime identifiers that have persisted SQLite stores in the default directory. + pub fn list_persisted_runtime_identifiers() -> Result, RuntimeError> { + let base = Self::default_directory(); + let Ok(entries) = std::fs::read_dir(&base) else { + return Ok(Vec::new()); + }; + + let mut runtime_identifiers = entries + .filter_map(|entry| entry.ok()) + .filter_map(|entry| entry.file_name().into_string().ok()) + .filter_map(|filename| decode_runtime_store_filename(&filename)) + .collect::>(); + runtime_identifiers.sort(); + runtime_identifiers.dedup(); + Ok(runtime_identifiers) + } + + /// Creates a SQLite runtime store at the provided path. + pub fn new(path: impl Into) -> Self { + Self { path: path.into() } + } + + /// Returns the SQLite database path for the store. + pub fn path(&self) -> &Path { + self.path.as_path() + } + + fn open(&self) -> Result { + if let Some(parent) = self.path.parent() { + std::fs::create_dir_all(parent) + .map_err(|error| RuntimeError::Store(error.to_string()))?; + } + let conn = Connection::open(&self.path).map_err(sqlite_error)?; + conn.busy_timeout(Duration::from_secs(5)) + .map_err(sqlite_error)?; + conn.pragma_update(None, "journal_mode", "WAL") + .map_err(sqlite_error)?; + conn.pragma_update(None, "foreign_keys", "ON") + .map_err(sqlite_error)?; + self.ensure_schema(&conn)?; + Ok(conn) + } + + fn ensure_schema(&self, conn: &Connection) -> Result<(), RuntimeError> { + conn.execute_batch( + r#" + CREATE TABLE IF NOT EXISTS agents ( + id TEXT PRIMARY KEY, + runtime_identifier TEXT NOT NULL, + name TEXT NOT NULL, + model TEXT NOT NULL, + provider_id TEXT NOT NULL, + team_dir TEXT NOT NULL, + tasks_namespace TEXT NOT NULL, + is_teammate INTEGER NOT NULL, + config_json TEXT NOT NULL, + hidden_tools_json TEXT NOT NULL, + max_rounds INTEGER, + teammate_identity_json TEXT, + rounds_since_task INTEGER NOT NULL, + idle_requested INTEGER NOT NULL, + status_json TEXT NOT NULL, + subagents_json TEXT NOT NULL, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS agent_memory ( + agent_id TEXT PRIMARY KEY, + revision INTEGER NOT NULL, + state_json TEXT NOT NULL, + updated_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS agent_runs ( + id TEXT PRIMARY KEY, + agent_id TEXT NOT NULL, + state TEXT NOT NULL, + error TEXT, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS tasks ( + namespace TEXT NOT NULL, + id INTEGER NOT NULL, + payload_json TEXT NOT NULL, + PRIMARY KEY (namespace, id) + ); + CREATE TABLE IF NOT EXISTS task_edges ( + namespace TEXT NOT NULL, + blocker_id INTEGER NOT NULL, + dependent_id INTEGER NOT NULL, + PRIMARY KEY (namespace, blocker_id, dependent_id) + ); + CREATE TABLE IF NOT EXISTS team_members ( + team_dir TEXT NOT NULL, + name TEXT NOT NULL, + summary_json TEXT NOT NULL, + PRIMARY KEY (team_dir, name) + ); + CREATE TABLE IF NOT EXISTS team_inbox ( + id TEXT PRIMARY KEY, + team_dir TEXT NOT NULL, + recipient TEXT NOT NULL, + payload_json TEXT NOT NULL, + delivery_state INTEGER NOT NULL, + created_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS team_requests ( + request_id TEXT PRIMARY KEY, + team_dir TEXT NOT NULL, + payload_json TEXT NOT NULL, + created_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS background_jobs ( + agent_id TEXT NOT NULL, + id TEXT NOT NULL, + payload_json TEXT NOT NULL, + notification_state INTEGER NOT NULL, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (agent_id, id) + ); + CREATE TABLE IF NOT EXISTS audit_events ( + id TEXT PRIMARY KEY, + scope TEXT NOT NULL, + event_type TEXT NOT NULL, + payload_json TEXT NOT NULL, + created_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS leases ( + key TEXT PRIMARY KEY, + owner TEXT NOT NULL, + expires_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS permission_rules ( + session_id TEXT NOT NULL, + project_id TEXT, + tool_name TEXT NOT NULL, + pattern TEXT, + allow INTEGER NOT NULL, + scope TEXT NOT NULL, + reason TEXT + ); + CREATE TABLE IF NOT EXISTS long_term_memory ( + record_id TEXT PRIMARY KEY, + agent_id TEXT NOT NULL, + kind TEXT NOT NULL, + content TEXT NOT NULL, + source_revision INTEGER NOT NULL, + created_at INTEGER NOT NULL, + metadata_json TEXT NOT NULL + ); + CREATE VIRTUAL TABLE IF NOT EXISTS long_term_memory_fts USING fts5( + record_id UNINDEXED, + agent_id UNINDEXED, + content + ); + CREATE TABLE IF NOT EXISTS long_term_memory_cursor ( + agent_id TEXT PRIMARY KEY, + cursor_json TEXT NOT NULL, + updated_at INTEGER NOT NULL + ); + "#, + ) + .map_err(sqlite_error)?; + self.migrate_background_jobs_schema(conn)?; + self.migrate_permission_rules_schema(conn) + } + + fn migrate_permission_rules_schema(&self, conn: &Connection) -> Result<(), RuntimeError> { + let Some(schema_sql) = conn + .query_row( + "SELECT sql FROM sqlite_master WHERE type = 'table' AND name = 'permission_rules'", + [], + |row| row.get::<_, String>(0), + ) + .optional() + .map_err(sqlite_error)? + else { + return Ok(()); + }; + + if !schema_sql.contains("project_id") { + conn.execute_batch("ALTER TABLE permission_rules ADD COLUMN project_id TEXT;") + .map_err(sqlite_error)?; + } + + // Rules written before a refusal could carry its reason have no such + // column; they gain a nullable one and keep loading as reasonless. + if !schema_sql.contains("reason") { + conn.execute_batch("ALTER TABLE permission_rules ADD COLUMN reason TEXT;") + .map_err(sqlite_error)?; + } + + // Ensure indexes exist (safe to run every time). + conn.execute_batch( + r#" + CREATE INDEX IF NOT EXISTS idx_perm_session ON permission_rules (session_id); + CREATE INDEX IF NOT EXISTS idx_perm_project ON permission_rules (project_id); + CREATE INDEX IF NOT EXISTS idx_perm_global ON permission_rules (scope); + "#, + ) + .map_err(sqlite_error)?; + + Ok(()) + } + + fn migrate_background_jobs_schema(&self, conn: &Connection) -> Result<(), RuntimeError> { + let Some(schema_sql) = conn + .query_row( + "SELECT sql FROM sqlite_master WHERE type = 'table' AND name = 'background_jobs'", + [], + |row| row.get::<_, String>(0), + ) + .optional() + .map_err(sqlite_error)? + else { + return Ok(()); + }; + + if schema_sql.contains("PRIMARY KEY (agent_id, id)") + || schema_sql.contains("PRIMARY KEY(agent_id, id)") + { + return Ok(()); + } + + conn.execute_batch( + r#" + ALTER TABLE background_jobs RENAME TO background_jobs_legacy; + CREATE TABLE background_jobs ( + agent_id TEXT NOT NULL, + id TEXT NOT NULL, + payload_json TEXT NOT NULL, + notification_state INTEGER NOT NULL, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (agent_id, id) + ); + INSERT INTO background_jobs (agent_id, id, payload_json, notification_state, created_at, updated_at) + SELECT agent_id, id, payload_json, notification_state, created_at, updated_at + FROM background_jobs_legacy; + DROP TABLE background_jobs_legacy; + "#, + ) + .map_err(sqlite_error)?; + + Ok(()) + } + + fn write_agent( + &self, + conn: &Connection, + record: &PersistedAgentRecord, + ) -> Result<(), RuntimeError> { + let now = now_secs(); + conn.execute( + r#" + INSERT INTO agents ( + id, runtime_identifier, name, model, provider_id, team_dir, tasks_namespace, is_teammate, config_json, + hidden_tools_json, max_rounds, teammate_identity_json, rounds_since_task, + idle_requested, status_json, subagents_json, created_at, updated_at + ) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18) + ON CONFLICT(id) DO UPDATE SET + runtime_identifier = excluded.runtime_identifier, + name = excluded.name, + model = excluded.model, + provider_id = excluded.provider_id, + team_dir = excluded.team_dir, + tasks_namespace = excluded.tasks_namespace, + is_teammate = excluded.is_teammate, + config_json = excluded.config_json, + hidden_tools_json = excluded.hidden_tools_json, + max_rounds = excluded.max_rounds, + teammate_identity_json = excluded.teammate_identity_json, + rounds_since_task = excluded.rounds_since_task, + idle_requested = excluded.idle_requested, + status_json = excluded.status_json, + subagents_json = excluded.subagents_json, + updated_at = excluded.updated_at + "#, + params![ + record.id, + record.runtime_identifier, + record.name, + record.model, + record.provider_id.as_str(), + record.config.team.team_dir.to_string_lossy().into_owned(), + record.config.task.tasks_dir.to_string_lossy().into_owned(), + i64::from(record.teammate_identity.is_some()), + to_json(&record.config)?, + to_json(&record.hidden_tools)?, + record.max_rounds.map(|value| value as i64), + maybe_json(&record.teammate_identity)?, + record.rounds_since_task as i64, + i64::from(record.idle_requested), + to_json(&record.status)?, + to_json(&record.subagents)?, + now, + now, + ], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn write_agent_memory( + &self, + conn: &Connection, + agent_id: &str, + memory: &AgentMemoryState, + ) -> Result<(), RuntimeError> { + conn.execute( + r#" + INSERT INTO agent_memory (agent_id, revision, state_json, updated_at) + VALUES (?1, ?2, ?3, ?4) + ON CONFLICT(agent_id) DO UPDATE SET + revision = excluded.revision, + state_json = excluded.state_json, + updated_at = excluded.updated_at + "#, + params![ + agent_id, + memory.revision as i64, + to_json(memory)?, + now_secs() + ], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn team_key(path: &Path) -> String { + path.to_string_lossy().into_owned() + } + + fn task_namespace(path: &Path) -> String { + path.to_string_lossy().into_owned() + } +} + +impl AgentStore for SqliteRuntimeStore { + fn prepare_recovery(&self) -> Result<(), RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + + tx.execute( + "UPDATE team_inbox SET delivery_state = ?1 WHERE delivery_state = ?2", + params![DELIVERY_PENDING, DELIVERY_INFLIGHT], + ) + .map_err(sqlite_error)?; + tx.execute( + "UPDATE background_jobs SET notification_state = ?1 WHERE notification_state = ?2", + params![DELIVERY_PENDING, DELIVERY_INFLIGHT], + ) + .map_err(sqlite_error)?; + + { + let mut stmt = tx + .prepare("SELECT agent_id, id, payload_json FROM background_jobs") + .map_err(sqlite_error)?; + let rows = stmt + .query_map([], |row| { + Ok(( + row.get::<_, String>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + )) + }) + .map_err(sqlite_error)?; + for row in rows { + let (agent_id, id, payload) = row.map_err(sqlite_error)?; + let mut task: BackgroundTaskSummary = from_json(&payload)?; + if task.status == BackgroundTaskStatus::Running { + task.status = BackgroundTaskStatus::Interrupted; + tx.execute( + "UPDATE background_jobs SET payload_json = ?3, notification_state = ?4, updated_at = ?5 WHERE agent_id = ?1 AND id = ?2", + params![agent_id, id, to_json(&task)?, DELIVERY_PENDING, now_secs()], + ) + .map_err(sqlite_error)?; + } + } + } + + tx.execute( + "DELETE FROM leases WHERE expires_at <= ?1", + params![now_secs()], + ) + .map_err(sqlite_error)?; + prune_stale_runtime_leases(&tx)?; + tx.commit().map_err(sqlite_error) + } + + fn create_agent( + &self, + record: &PersistedAgentRecord, + memory: &AgentMemoryState, + ) -> Result<(), RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + self.write_agent(&tx, record)?; + self.write_agent_memory(&tx, &record.id, memory)?; + tx.commit().map_err(sqlite_error) + } + + fn save_agent_record(&self, record: &PersistedAgentRecord) -> Result<(), RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + self.write_agent(&tx, record)?; + tx.commit().map_err(sqlite_error) + } + + fn save_agent_memory( + &self, + agent_id: &str, + memory: &AgentMemoryState, + ) -> Result<(), RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + self.write_agent_memory(&tx, agent_id, memory)?; + tx.commit().map_err(sqlite_error) + } + + fn load_agent(&self, agent_id: &str) -> Result, RuntimeError> { + let conn = self.open()?; + let record = conn + .query_row( + r#" + SELECT + id, runtime_identifier, name, model, provider_id, config_json, + hidden_tools_json, max_rounds, teammate_identity_json, rounds_since_task, + idle_requested, status_json, subagents_json + FROM agents WHERE id = ?1 + "#, + params![agent_id], + |row| { + let provider_id: String = row.get(4)?; + let config_json: String = row.get(5)?; + let hidden_tools_json: String = row.get(6)?; + let teammate_identity_json: Option = row.get(8)?; + let status_json: String = row.get(11)?; + let subagents_json: String = row.get(12)?; + Ok(PersistedAgentRecord { + id: row.get(0)?, + runtime_identifier: row.get(1)?, + name: row.get(2)?, + model: row.get(3)?, + provider_id: ProviderId::from(provider_id), + config: from_json(&config_json).map_err(to_sql_error)?, + hidden_tools: from_json(&hidden_tools_json).map_err(to_sql_error)?, + max_rounds: row.get::<_, Option>(7)?.map(|value| value as usize), + teammate_identity: teammate_identity_json + .map(|json| from_json(&json)) + .transpose() + .map_err(to_sql_error)?, + rounds_since_task: row.get::<_, i64>(9)? as usize, + idle_requested: row.get::<_, i64>(10)? != 0, + status: from_json(&status_json).map_err(to_sql_error)?, + subagents: from_json(&subagents_json).map_err(to_sql_error)?, + }) + }, + ) + .optional() + .map_err(sqlite_error)?; + let Some(record) = record else { + return Ok(None); + }; + + let memory = conn + .query_row( + "SELECT state_json FROM agent_memory WHERE agent_id = ?1", + params![agent_id], + |row| { + let state_json: String = row.get(0)?; + from_json(&state_json).map_err(to_sql_error) + }, + ) + .optional() + .map_err(sqlite_error)?; + let Some(memory) = memory else { + return Err(RuntimeError::Store(format!( + "Agent '{agent_id}' is missing persisted memory" + ))); + }; + + Ok(Some(LoadedAgentState { record, memory })) + } + + fn list_agents(&self) -> Result, RuntimeError> { + let conn = self.open()?; + let mut stmt = conn + .prepare("SELECT id FROM agents ORDER BY created_at, id") + .map_err(sqlite_error)?; + let ids = stmt + .query_map([], |row| row.get::<_, String>(0)) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error)?; + ids.into_iter() + .map(|id| { + self.load_agent(&id)? + .ok_or_else(|| RuntimeError::Store(format!("Agent '{id}' disappeared"))) + }) + .collect() + } + + fn list_agents_by_runtime( + &self, + runtime_identifier: &str, + ) -> Result, RuntimeError> { + let conn = self.open()?; + let mut stmt = conn + .prepare("SELECT id FROM agents WHERE runtime_identifier = ?1 ORDER BY created_at, id") + .map_err(sqlite_error)?; + let ids = stmt + .query_map(params![runtime_identifier], |row| row.get::<_, String>(0)) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error)?; + ids.into_iter() + .map(|id| { + self.load_agent(&id)? + .ok_or_else(|| RuntimeError::Store(format!("Agent '{id}' disappeared"))) + }) + .collect() + } +} + +impl RunStore for SqliteRuntimeStore { + fn start_run(&self, agent_id: &str) -> Result { + let run_id = next_id("run"); + let conn = self.open()?; + conn.execute( + "INSERT INTO agent_runs (id, agent_id, state, error, created_at, updated_at) VALUES (?1, ?2, 'running', NULL, ?3, ?3)", + params![run_id, agent_id, now_secs()], + ) + .map_err(sqlite_error)?; + Ok(run_id) + } + + fn update_run_state( + &self, + run_id: &str, + state: &str, + error: Option<&str>, + ) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "UPDATE agent_runs SET state = ?2, error = ?3, updated_at = ?4 WHERE id = ?1", + params![run_id, state, error, now_secs()], + ) + .map_err(sqlite_error)?; + Ok(()) + } + + fn finish_run(&self, run_id: &str) -> Result<(), RuntimeError> { + self.update_run_state(run_id, "finished", None) + } + + fn fail_run(&self, run_id: &str, error: &str) -> Result<(), RuntimeError> { + self.update_run_state(run_id, "failed", Some(error)) + } +} + +impl TaskStore for SqliteRuntimeStore { + fn load_tasks(&self, namespace: &Path) -> Result, RuntimeError> { + let conn = self.open()?; + self.load_tasks_from_conn(&conn, namespace) + } + + fn capture_tasks(&self, namespace: &Path) -> Result { + Ok(TaskStateSnapshot { + tasks: self.load_tasks(namespace)?, + }) + } + + fn restore_tasks( + &self, + namespace: &Path, + snapshot: &TaskStateSnapshot, + ) -> Result<(), RuntimeError> { + self.replace_tasks(namespace, &snapshot.tasks) + } + + fn replace_tasks(&self, namespace: &Path, tasks: &[TaskItem]) -> Result<(), RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + self.replace_tasks_in_conn(&tx, namespace, tasks)?; + tx.commit().map_err(sqlite_error) + } + + fn mutate( + &self, + namespace: &Path, + mutation: &mut dyn FnMut(&mut Vec) -> Result<(), RuntimeError>, + ) -> Result<(), RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + let mut tasks = self.load_tasks_from_conn(&tx, namespace)?; + mutation(&mut tasks)?; + self.replace_tasks_in_conn(&tx, namespace, &tasks)?; + tx.commit().map_err(sqlite_error) + } +} + +impl AuditStore for SqliteRuntimeStore { + fn record_audit_event( + &self, + scope: &str, + event_type: &str, + payload: serde_json::Value, + ) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "INSERT INTO audit_events (id, scope, event_type, payload_json, created_at) VALUES (?1, ?2, ?3, ?4, ?5)", + params![next_id("audit"), scope, event_type, payload.to_string(), now_secs()], + ) + .map_err(sqlite_error)?; + Ok(()) + } +} + +impl LeaseStore for SqliteRuntimeStore { + fn acquire_lease(&self, key: &str, owner: &str, ttl: Duration) -> Result { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + let now = now_secs(); + tx.execute("DELETE FROM leases WHERE expires_at <= ?1", params![now]) + .map_err(sqlite_error)?; + prune_stale_runtime_leases(&tx)?; + let inserted = tx + .execute( + "INSERT OR IGNORE INTO leases (key, owner, expires_at) VALUES (?1, ?2, ?3)", + params![key, owner, now + ttl.as_secs() as i64], + ) + .map_err(sqlite_error)?; + tx.commit().map_err(sqlite_error)?; + Ok(inserted == 1) + } + + fn release_lease(&self, key: &str, owner: &str) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "DELETE FROM leases WHERE key = ?1 AND owner = ?2", + params![key, owner], + ) + .map_err(sqlite_error)?; + Ok(()) + } +} + +impl PermissionRuleStore for SqliteRuntimeStore { + fn save_rules( + &self, + session_id: &str, + project_id: Option<&str>, + rules: &[RememberedRule], + ) -> Result<(), RuntimeError> { + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + + // Only delete session-scoped rules for this session; project and global + // rules are managed separately and must not be removed here. + let session_scope = to_json(&PermissionRuleScope::Session)?; + tx.execute( + "DELETE FROM permission_rules WHERE session_id = ?1 AND scope = ?2", + params![session_id, session_scope], + ) + .map_err(sqlite_error)?; + + for rule in rules { + let scope_str = to_json(&rule.scope)?; + tx.execute( + r#" + INSERT INTO permission_rules (session_id, project_id, tool_name, pattern, allow, scope, reason) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7) + "#, + params![ + session_id, + project_id, + rule.key.tool_name, + rule.key.pattern, + rule.allow as i32, + scope_str, + rule.reason, + ], + ) + .map_err(sqlite_error)?; + } + + tx.commit().map_err(sqlite_error)?; + Ok(()) + } + + fn load_rules( + &self, + session_id: &str, + project_id: Option<&str>, + ) -> Result, RuntimeError> { + let conn = self.open()?; + + let session_scope = to_json(&PermissionRuleScope::Session)?; + let project_scope = to_json(&PermissionRuleScope::Project)?; + let global_scope = to_json(&PermissionRuleScope::Global)?; + + // UNION of session-scoped, project-scoped (if project_id provided), + // and global-scoped rules. + let sql = r#" + SELECT tool_name, pattern, allow, scope, reason + FROM permission_rules + WHERE session_id = ?1 AND scope = ?2 + UNION + SELECT tool_name, pattern, allow, scope, reason + FROM permission_rules + WHERE project_id IS NOT NULL AND project_id = ?3 AND scope = ?4 + UNION + SELECT tool_name, pattern, allow, scope, reason + FROM permission_rules + WHERE scope = ?5 + "#; + + let mut stmt = conn.prepare(sql).map_err(sqlite_error)?; + + // When project_id is None we pass an empty string; the IS NOT NULL guard + // in the project clause prevents accidental matches. + let project_id_param = project_id.unwrap_or(""); + + let rows = stmt + .query_map( + params![ + session_id, + session_scope, + project_id_param, + project_scope, + global_scope, + ], + |row| { + Ok(( + row.get::<_, String>(0)?, + row.get::<_, Option>(1)?, + row.get::<_, i32>(2)?, + row.get::<_, String>(3)?, + row.get::<_, Option>(4)?, + )) + }, + ) + .map_err(sqlite_error)?; + + let mut rules = Vec::new(); + for row in rows { + let (tool_name, pattern, allow, scope_str, reason) = row.map_err(sqlite_error)?; + let scope: PermissionRuleScope = from_json(&scope_str)?; + rules.push(RememberedRule { + key: RuleKey { tool_name, pattern }, + allow: allow != 0, + scope, + reason, + }); + } + Ok(rules) + } + + fn clear_rules(&self, session_id: &str) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + "DELETE FROM permission_rules WHERE session_id = ?1", + params![session_id], + ) + .map_err(sqlite_error)?; + Ok(()) + } +} + +impl MemoryStore for SqliteRuntimeStore { + fn upsert_records(&self, records: &[MemoryRecord]) -> Result<(), RuntimeError> { + if records.is_empty() { + return Ok(()); + } + + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + + for record in records { + tx.execute( + r#" + INSERT INTO long_term_memory ( + record_id, agent_id, kind, content, source_revision, created_at, metadata_json + ) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7) + ON CONFLICT(record_id) DO UPDATE SET + agent_id = excluded.agent_id, + kind = excluded.kind, + content = excluded.content, + source_revision = excluded.source_revision, + created_at = excluded.created_at, + metadata_json = excluded.metadata_json + "#, + params![ + record.record_id, + record.agent_id, + format!("{:?}", record.kind).to_lowercase(), + record.content, + record.source_revision as i64, + record.created_at, + record.metadata_json, + ], + ) + .map_err(sqlite_error)?; + tx.execute( + "DELETE FROM long_term_memory_fts WHERE record_id = ?1", + params![record.record_id], + ) + .map_err(sqlite_error)?; + tx.execute( + "INSERT INTO long_term_memory_fts (record_id, agent_id, content) VALUES (?1, ?2, ?3)", + params![record.record_id, record.agent_id, record.content], + ) + .map_err(sqlite_error)?; + } + + tx.commit().map_err(sqlite_error) + } + + fn search_records_with_options( + &self, + request: &MemorySearchRequest, + ) -> Result, RuntimeError> { + if request.query.trim().is_empty() || request.limit == 0 { + return Ok(Vec::new()); + } + if request.filter.pinned == Some(true) || request.filter.source.is_some() { + return Ok(Vec::new()); + } + let Some(query) = fts_query(&request.query) else { + return Ok(Vec::new()); + }; + + let conn = self.open()?; + let kind = request + .filter + .kind + .map(|kind| format!("{kind:?}").to_lowercase()); + let mut stmt = conn + .prepare( + r#" + SELECT + memory.record_id, + memory.agent_id, + memory.kind, + memory.content, + memory.source_revision, + memory.created_at, + memory.metadata_json, + bm25(long_term_memory_fts) AS rank + FROM long_term_memory_fts + JOIN long_term_memory AS memory ON memory.record_id = long_term_memory_fts.record_id + WHERE long_term_memory_fts.agent_id = ?1 + AND long_term_memory_fts.content MATCH ?2 + AND (?3 IS NULL OR memory.kind = ?3) + AND (?4 IS NULL OR memory.created_at >= ?4) + AND (?5 IS NULL OR memory.created_at <= ?5) + ORDER BY rank, memory.created_at DESC + LIMIT ?6 + "#, + ) + .map_err(sqlite_error)?; + + stmt.query_map( + params![ + request.agent_id, + query, + kind, + request.filter.created_from, + request.filter.created_to, + request.limit as i64, + ], + |row| { + let kind = row.get::<_, String>(2)?; + Ok(MemoryRecord { + record_id: row.get(0)?, + agent_id: row.get(1)?, + kind: parse_memory_kind(&kind), + content: row.get(3)?, + source_revision: row.get::<_, i64>(4)? as u64, + created_at: row.get(5)?, + metadata_json: row.get(6)?, + source: None, + pinned: false, + score: row.get::<_, Option>(7)?, + }) + }, + ) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error) + } + + fn list_records(&self, request: &MemoryListRequest) -> Result { + let limit = request.limit.min(crate::memory::MAX_MEMORY_LIST_PAGE_SIZE); + if limit == 0 || request.filter.pinned == Some(true) || request.filter.source.is_some() { + return Ok(MemoryListPage { + records: Vec::new(), + next_cursor: None, + }); + } + let conn = self.open()?; + let kind = request + .filter + .kind + .map(|kind| format!("{kind:?}").to_lowercase()); + let cursor_time = request.cursor.as_ref().map(|cursor| cursor.created_at); + let cursor_id = request + .cursor + .as_ref() + .map(|cursor| cursor.record_id.as_str()); + let direction = match request.sort { + MemoryListSort::Newest => "<", + MemoryListSort::Oldest => ">", + }; + let order = match request.sort { + MemoryListSort::Newest => "DESC", + MemoryListSort::Oldest => "ASC", + }; + let sql = format!( + r#" + SELECT record_id, agent_id, kind, content, source_revision, created_at, metadata_json + FROM long_term_memory + WHERE agent_id = ?1 + AND (?2 IS NULL OR kind = ?2) + AND (?3 IS NULL OR created_at >= ?3) + AND (?4 IS NULL OR created_at <= ?4) + AND (?5 IS NULL OR created_at {direction} ?5 + OR (created_at = ?5 AND record_id {direction} ?6)) + ORDER BY created_at {order}, record_id {order} + LIMIT ?7 + "# + ); + let mut stmt = conn.prepare(&sql).map_err(sqlite_error)?; + let mut records = stmt + .query_map( + params![ + request.agent_id, + kind, + request.filter.created_from, + request.filter.created_to, + cursor_time, + cursor_id, + limit.saturating_add(1) as i64, + ], + |row| { + let kind = row.get::<_, String>(2)?; + Ok(MemoryRecord { + record_id: row.get(0)?, + agent_id: row.get(1)?, + kind: parse_memory_kind(&kind), + content: row.get(3)?, + source_revision: row.get::<_, i64>(4)? as u64, + created_at: row.get(5)?, + metadata_json: row.get(6)?, + source: None, + pinned: false, + score: None, + }) + }, + ) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error)?; + let has_more = records.len() > limit; + records.truncate(limit); + let next_cursor = if has_more { + records.last().map(|record| MemoryListCursor { + created_at: record.created_at, + record_id: record.record_id.clone(), + }) + } else { + None + }; + Ok(MemoryListPage { + records, + next_cursor, + }) + } + + fn get_record( + &self, + agent_id: &str, + record_id: &str, + ) -> Result, RuntimeError> { + let conn = self.open()?; + conn.query_row( + r#" + SELECT record_id, agent_id, kind, content, source_revision, created_at, metadata_json + FROM long_term_memory WHERE agent_id = ?1 AND record_id = ?2 + "#, + params![agent_id, record_id], + |row| { + let kind = row.get::<_, String>(2)?; + Ok(MemoryRecord { + record_id: row.get(0)?, + agent_id: row.get(1)?, + kind: parse_memory_kind(&kind), + content: row.get(3)?, + source_revision: row.get::<_, i64>(4)? as u64, + created_at: row.get(5)?, + metadata_json: row.get(6)?, + source: None, + pinned: false, + score: None, + }) + }, + ) + .optional() + .map_err(sqlite_error) + } + + fn count_records(&self, agent_id: &str) -> Result { + let conn = self.open()?; + conn.query_row( + "SELECT COUNT(*) FROM long_term_memory WHERE agent_id = ?1", + params![agent_id], + |row| row.get::<_, i64>(0), + ) + .map(|count| count as usize) + .map_err(sqlite_error) + } + + fn delete_records(&self, record_ids: &[String]) -> Result<(), RuntimeError> { + if record_ids.is_empty() { + return Ok(()); + } + + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + for record_id in record_ids { + tx.execute( + "DELETE FROM long_term_memory_fts WHERE record_id = ?1", + params![record_id], + ) + .map_err(sqlite_error)?; + tx.execute( + "DELETE FROM long_term_memory WHERE record_id = ?1", + params![record_id], + ) + .map_err(sqlite_error)?; + } + tx.commit().map_err(sqlite_error) + } + + fn tombstone_records( + &self, + agent_id: &str, + record_ids: &[String], + ) -> Result { + if record_ids.is_empty() { + return Ok(0); + } + + let mut conn = self.open()?; + let tx = conn + .transaction_with_behavior(TransactionBehavior::Immediate) + .map_err(sqlite_error)?; + let mut affected = 0usize; + for record_id in record_ids { + tx.execute( + "DELETE FROM long_term_memory_fts WHERE record_id = ?1", + params![record_id], + ) + .map_err(sqlite_error)?; + affected += tx + .execute( + "DELETE FROM long_term_memory WHERE record_id = ?1 AND agent_id = ?2", + params![record_id, agent_id], + ) + .map_err(sqlite_error)?; + } + tx.commit().map_err(sqlite_error)?; + Ok(affected) + } + + fn load_agent_memory_cursor( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + let conn = self.open()?; + conn.query_row( + "SELECT cursor_json FROM long_term_memory_cursor WHERE agent_id = ?1", + params![agent_id], + |row| row.get::<_, String>(0), + ) + .optional() + .map_err(sqlite_error)? + .map(|json| from_json(&json)) + .transpose() + } + + fn save_agent_memory_cursor( + &self, + agent_id: &str, + cursor: &MemoryCursor, + ) -> Result<(), RuntimeError> { + let conn = self.open()?; + conn.execute( + r#" + INSERT INTO long_term_memory_cursor (agent_id, cursor_json, updated_at) + VALUES (?1, ?2, ?3) + ON CONFLICT(agent_id) DO UPDATE SET + cursor_json = excluded.cursor_json, + updated_at = excluded.updated_at + "#, + params![agent_id, to_json(cursor)?, now_secs()], + ) + .map_err(sqlite_error)?; + Ok(()) + } +} + +impl SqliteRuntimeStore { + fn load_tasks_from_conn( + &self, + conn: &Connection, + namespace: &Path, + ) -> Result, RuntimeError> { + let mut stmt = conn + .prepare("SELECT payload_json FROM tasks WHERE namespace = ?1 ORDER BY id") + .map_err(sqlite_error)?; + let rows = stmt + .query_map(params![Self::task_namespace(namespace)], |row| { + row.get::<_, String>(0) + }) + .map_err(sqlite_error)?; + let mut tasks = Vec::new(); + for row in rows { + tasks.push(from_json(&row.map_err(sqlite_error)?)?); + } + Ok(tasks) + } + + fn replace_tasks_in_conn( + &self, + conn: &Connection, + namespace: &Path, + tasks: &[TaskItem], + ) -> Result<(), RuntimeError> { + let namespace = Self::task_namespace(namespace); + conn.execute( + "DELETE FROM tasks WHERE namespace = ?1", + params![namespace.clone()], + ) + .map_err(sqlite_error)?; + conn.execute( + "DELETE FROM task_edges WHERE namespace = ?1", + params![namespace.clone()], + ) + .map_err(sqlite_error)?; + for task in tasks { + conn.execute( + "INSERT INTO tasks (namespace, id, payload_json) VALUES (?1, ?2, ?3)", + params![namespace.clone(), task.id as i64, to_json(task)?], + ) + .map_err(sqlite_error)?; + for blocker in &task.blocked_by { + conn.execute( + "INSERT OR IGNORE INTO task_edges (namespace, blocker_id, dependent_id) VALUES (?1, ?2, ?3)", + params![namespace.clone(), *blocker as i64, task.id as i64], + ) + .map_err(sqlite_error)?; + } + } + Ok(()) + } +} + +fn to_json(value: &T) -> Result { + serde_json::to_string(value).map_err(|error| RuntimeError::Store(error.to_string())) +} + +fn maybe_json(value: &Option) -> Result, RuntimeError> { + value.as_ref().map(to_json).transpose() +} + +fn from_json(value: &str) -> Result { + serde_json::from_str(value).map_err(|error| RuntimeError::Store(error.to_string())) +} + +fn sqlite_error(error: rusqlite::Error) -> RuntimeError { + RuntimeError::Store(error.to_string()) +} + +fn parse_memory_kind(kind: &str) -> crate::memory::MemoryRecordKind { + match kind { + "summary" => crate::memory::MemoryRecordKind::Summary, + "fact" => crate::memory::MemoryRecordKind::Fact, + _ => crate::memory::MemoryRecordKind::Episode, + } +} + +fn fts_query(query: &str) -> Option { + let tokens = query + .split(|ch: char| !ch.is_alphanumeric()) + .filter(|token| !token.is_empty()) + .map(|token| format!("\"{token}\"")) + .collect::>(); + + if tokens.is_empty() { + None + } else { + Some(tokens.join(" OR ")) + } +} + +fn to_sql_error(error: RuntimeError) -> rusqlite::Error { + rusqlite::Error::FromSqlConversionFailure( + 0, + rusqlite::types::Type::Text, + Box::new(std::io::Error::other(error.to_string())), + ) +} + +fn next_id(prefix: &str) -> String { + let counter = NEXT_STORE_ID.fetch_add(1, Ordering::Relaxed); + format!("{prefix}-{:x}-{:x}", now_nanos(), counter) +} + +fn now_secs() -> i64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64 +} + +fn now_nanos() -> u128 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() +} + +fn prune_stale_runtime_leases(tx: &rusqlite::Transaction<'_>) -> Result<(), RuntimeError> { + let mut stmt = tx + .prepare("SELECT key, owner FROM leases") + .map_err(sqlite_error)?; + let leases = stmt + .query_map([], |row| { + Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)) + }) + .map_err(sqlite_error)? + .collect::, _>>() + .map_err(sqlite_error)?; + drop(stmt); + + for (key, owner) in leases { + if runtime_owner_is_stale(&owner) { + tx.execute("DELETE FROM leases WHERE key = ?1", params![key]) + .map_err(sqlite_error)?; + } + } + + Ok(()) +} + +fn runtime_owner_is_stale(owner: &str) -> bool { + let Some(pid) = owner + .strip_prefix("runtime-") + .and_then(|value| value.parse::().ok()) + else { + return false; + }; + + #[cfg(unix)] + { + let pid = pid as i32; + let result = unsafe { libc::kill(pid, 0) }; + if result == 0 { + return false; + } + + match std::io::Error::last_os_error().raw_os_error() { + Some(code) if code == libc::ESRCH => true, + Some(code) if code == libc::EPERM => false, + _ => false, + } + } + + #[cfg(windows)] + { + use windows_sys::Win32::{ + Foundation::{CloseHandle, STILL_ACTIVE}, + System::Threading::{ + GetExitCodeProcess, OpenProcess, PROCESS_QUERY_LIMITED_INFORMATION, + }, + }; + + const ERROR_ACCESS_DENIED: i32 = 5; + const ERROR_INVALID_PARAMETER: i32 = 87; + + unsafe { + let handle = OpenProcess(PROCESS_QUERY_LIMITED_INFORMATION, 0, pid); + if handle.is_null() { + return match std::io::Error::last_os_error().raw_os_error() { + Some(ERROR_INVALID_PARAMETER) => true, + Some(ERROR_ACCESS_DENIED) => false, + _ => false, + }; + } + + let mut exit_code = 0u32; + let result = GetExitCodeProcess(handle, &mut exit_code); + let close_result = CloseHandle(handle); + debug_assert_ne!(close_result, 0, "process handle should close"); + + if result == 0 { + return false; + } + + exit_code != STILL_ACTIVE as u32 + } + } + + #[cfg(not(any(unix, windows)))] + { + false + } +} + +fn encode_runtime_identifier(runtime_identifier: &str) -> String { + let mut encoded = String::with_capacity(runtime_identifier.len() * 2); + for byte in runtime_identifier.as_bytes() { + use std::fmt::Write as _; + let _ = write!(&mut encoded, "{byte:02x}"); + } + encoded +} + +fn decode_runtime_store_filename(filename: &str) -> Option { + let encoded = filename.strip_prefix("runtime-")?.strip_suffix(".sqlite")?; + if encoded.len() % 2 != 0 || encoded.is_empty() { + return None; + } + + let mut bytes = Vec::with_capacity(encoded.len() / 2); + let mut index = 0; + while index < encoded.len() { + let byte = u8::from_str_radix(&encoded[index..index + 2], 16).ok()?; + bytes.push(byte); + index += 2; + } + String::from_utf8(bytes).ok() +} + +#[cfg(not(test))] +fn default_store_dir() -> PathBuf { + crate::default_paths::workspace_default_paths().root_dir +} + +#[cfg(test)] +thread_local! { + /// Every default store directory handed out on the current thread. + /// + /// Each one is unique, so recording them lets a test name the database its + /// own builder would have used — which is the only way to assert that a + /// builder given an explicit store left the default alone. + static DEFAULT_STORE_DIRS: std::cell::RefCell> = + const { std::cell::RefCell::new(Vec::new()) }; +} + +#[cfg(test)] +fn default_store_dir() -> PathBuf { + let suffix = NEXT_TEST_STORE_ID.fetch_add(1, Ordering::Relaxed); + let dir = std::env::temp_dir() + .join("mentra-test-runtime") + .join(format!("process-{}-{suffix}", std::process::id())); + DEFAULT_STORE_DIRS.with(|dirs| dirs.borrow_mut().push(dir.clone())); + dir +} + +/// The database paths of every default store constructed on this thread. +/// +/// Only `open()` creates the directory holding one, so an untouched default +/// store leaves nothing at these paths. +#[cfg(test)] +pub(crate) fn default_store_paths_on_this_thread() -> Vec { + DEFAULT_STORE_DIRS.with(|dirs| { + dirs.borrow() + .iter() + .map(|dir| dir.join("runtime.sqlite")) + .collect() + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::memory::{MemoryRecord, MemoryRecordKind, MemoryStore}; + + #[test] + fn runtime_identifier_round_trips_through_filename_encoding() { + let runtime_identifier = "chat/example 01"; + let filename = format!( + "runtime-{}.sqlite", + encode_runtime_identifier(runtime_identifier) + ); + assert_eq!( + decode_runtime_store_filename(&filename).as_deref(), + Some(runtime_identifier) + ); + } + + #[test] + fn path_for_runtime_identifier_uses_runtime_specific_filename() { + let path = SqliteRuntimeStore::path_for_runtime_identifier("session-a"); + assert!( + path.file_name() + .and_then(|name| name.to_str()) + .is_some_and(|name| name.starts_with("runtime-")) + ); + assert!( + path.file_name() + .and_then(|name| name.to_str()) + .is_some_and(|name| name.ends_with(".sqlite")) + ); + } + + #[test] + fn stale_runtime_owner_can_be_reclaimed() { + let store = SqliteRuntimeStore::new( + std::env::temp_dir().join(format!("mentra-store-lease-{}.sqlite", now_nanos())), + ); + let conn = Connection::open(store.path()).expect("open store"); + store.ensure_schema(&conn).expect("ensure schema"); + conn.execute( + "INSERT INTO leases (key, owner, expires_at) VALUES (?1, ?2, ?3)", + params!["agent:test", "runtime-999999", now_secs() + 3600], + ) + .expect("insert stale lease"); + + let acquired = store + .acquire_lease("agent:test", "runtime-123", Duration::from_secs(60)) + .expect("acquire lease"); + assert!(acquired); + } + + #[test] + fn background_tasks_are_scoped_per_agent() { + let store = SqliteRuntimeStore::new( + std::env::temp_dir().join(format!("mentra-store-background-{}.sqlite", now_nanos())), + ); + + store + .upsert_background_task( + "agent-a", + &BackgroundTaskSummary { + id: "bg-1".to_string(), + command: "echo a".to_string(), + cwd: std::env::temp_dir().join("a"), + status: BackgroundTaskStatus::Running, + output_preview: None, + }, + DELIVERY_ACKED, + ) + .expect("seed agent a background task"); + store + .upsert_background_task( + "agent-b", + &BackgroundTaskSummary { + id: "bg-1".to_string(), + command: "echo b".to_string(), + cwd: std::env::temp_dir().join("b"), + status: BackgroundTaskStatus::Finished, + output_preview: Some("done".to_string()), + }, + DELIVERY_PENDING, + ) + .expect("seed agent b background task"); + + let agent_a_tasks = store + .load_background_tasks("agent-a") + .expect("load agent a background tasks"); + let agent_b_tasks = store + .load_background_tasks("agent-b") + .expect("load agent b background tasks"); + + assert_eq!(agent_a_tasks.len(), 1); + assert_eq!(agent_a_tasks[0].command, "echo a"); + assert_eq!(agent_a_tasks[0].status, BackgroundTaskStatus::Running); + assert_eq!(agent_b_tasks.len(), 1); + assert_eq!(agent_b_tasks[0].command, "echo b"); + assert_eq!(agent_b_tasks[0].status, BackgroundTaskStatus::Finished); + } + + #[test] + fn fts_query_returns_none_when_input_has_no_searchable_terms() { + assert_eq!(fts_query("... --- \"\""), None); + } + + #[test] + fn sqlite_memory_search_sanitizes_punctuation_heavy_queries() { + let store = SqliteRuntimeStore::new( + std::env::temp_dir().join(format!("mentra-store-memory-{}.sqlite", now_nanos())), + ); + store + .upsert_records(&[MemoryRecord { + record_id: "episode:agent:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Episode, + content: "shared phrase alpha".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("seed".to_string()), + pinned: false, + score: None, + }]) + .expect("seed records"); + + let records = store + .search_records("agent-1", "(shared) alpha!!!", 10) + .expect("search records"); + assert_eq!(records.len(), 1); + assert_eq!(records[0].record_id, "episode:agent:1"); + } + + #[test] + fn sqlite_memory_search_ignores_non_searchable_queries() { + let store = SqliteRuntimeStore::new( + std::env::temp_dir().join(format!("mentra-store-empty-query-{}.sqlite", now_nanos())), + ); + store + .upsert_records(&[MemoryRecord { + record_id: "episode:agent:1".to_string(), + agent_id: "agent-1".to_string(), + kind: MemoryRecordKind::Episode, + content: "shared phrase alpha".to_string(), + source_revision: 1, + created_at: 1, + metadata_json: "{}".to_string(), + source: Some("seed".to_string()), + pinned: false, + score: None, + }]) + .expect("seed records"); + + let records = store + .search_records("agent-1", "... ---", 10) + .expect("search records"); + assert!(records.is_empty()); + } + + // -- PermissionRuleStore -- + + fn permission_store() -> SqliteRuntimeStore { + SqliteRuntimeStore::new( + std::env::temp_dir().join(format!("mentra-store-perm-{}.sqlite", now_nanos())), + ) + } + + #[test] + fn permission_rules_save_and_load_round_trip() { + use crate::session::PermissionRuleScope; + + let store = permission_store(); + + // Session-scoped rule under session-1 (no project). + let session_rule = RememberedRule { + key: RuleKey { + tool_name: "shell".to_string(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }; + // Project-scoped rule saved under session-1 with an explicit project_id. + let project_rule = RememberedRule { + key: RuleKey { + tool_name: "read".to_string(), + pattern: Some("/tmp/*".to_string()), + }, + allow: false, + scope: PermissionRuleScope::Project, + reason: None, + }; + + store + .save_rules("session-1", Some("proj-x"), &[session_rule, project_rule]) + .expect("save rules"); + + // Load with matching project_id: both session + project rules come back. + let loaded = store + .load_rules("session-1", Some("proj-x")) + .expect("load rules"); + + assert_eq!(loaded.len(), 2, "expected 2 rules, got {loaded:?}"); + + let shell_rule = loaded + .iter() + .find(|r| r.key.tool_name == "shell") + .expect("shell rule present"); + assert!(shell_rule.allow); + assert_eq!(shell_rule.scope, PermissionRuleScope::Session); + assert_eq!(shell_rule.key.pattern, None); + + let read_rule = loaded + .iter() + .find(|r| r.key.tool_name == "read") + .expect("read rule present"); + assert!(!read_rule.allow); + assert_eq!(read_rule.scope, PermissionRuleScope::Project); + assert_eq!(read_rule.key.pattern, Some("/tmp/*".to_string())); + } + + #[test] + fn permission_rules_clear_removes_all_for_session() { + use crate::session::PermissionRuleScope; + + let store = permission_store(); + let rules = vec![RememberedRule { + key: RuleKey { + tool_name: "shell".to_string(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }]; + + store + .save_rules("session-1", None, &rules) + .expect("save rules"); + store.clear_rules("session-1").expect("clear rules"); + + let loaded = store + .load_rules("session-1", None) + .expect("load rules after clear"); + assert!(loaded.is_empty()); + } + + #[test] + fn permission_rules_are_scoped_per_session() { + use crate::session::PermissionRuleScope; + + // Use an isolated store so global rules from one session don't bleed + // into assertions about another session. + let store = permission_store(); + let rules_a = vec![RememberedRule { + key: RuleKey { + tool_name: "shell".to_string(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }]; + // session-b uses a project-scoped rule (not global) so it doesn't show + // up when loading session-a without a matching project_id. + let rules_b = vec![RememberedRule { + key: RuleKey { + tool_name: "read".to_string(), + pattern: None, + }, + allow: false, + scope: PermissionRuleScope::Project, + reason: None, + }]; + + store + .save_rules("session-a", None, &rules_a) + .expect("save rules a"); + store + .save_rules("session-b", Some("proj-b"), &rules_b) + .expect("save rules b"); + + // Load session-a without project_id: only its own session-scoped rules. + let loaded_a = store.load_rules("session-a", None).expect("load rules a"); + // Load session-b with its project_id: project-scoped rules come back. + let loaded_b = store + .load_rules("session-b", Some("proj-b")) + .expect("load rules b"); + + assert_eq!(loaded_a.len(), 1, "session-a should have 1 rule"); + assert_eq!(loaded_a[0].key.tool_name, "shell"); + assert_eq!(loaded_b.len(), 1, "session-b should have 1 rule"); + assert_eq!(loaded_b[0].key.tool_name, "read"); + } + + #[test] + fn permission_rules_save_replaces_existing_session_rules() { + use crate::session::PermissionRuleScope; + + let store = permission_store(); + let initial = vec![RememberedRule { + key: RuleKey { + tool_name: "shell".to_string(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }]; + + store + .save_rules("session-1", None, &initial) + .expect("save initial"); + + // Replace with a different session-scoped rule. + let updated = vec![RememberedRule { + key: RuleKey { + tool_name: "write".to_string(), + pattern: None, + }, + allow: false, + scope: PermissionRuleScope::Session, + reason: None, + }]; + + store + .save_rules("session-1", None, &updated) + .expect("save updated"); + + let loaded = store + .load_rules("session-1", None) + .expect("load after replace"); + assert_eq!(loaded.len(), 1, "should have exactly 1 rule after replace"); + assert_eq!(loaded[0].key.tool_name, "write"); + assert!(!loaded[0].allow); + } + + #[test] + fn permission_rules_load_returns_empty_for_unknown_session() { + let store = permission_store(); + let loaded = store + .load_rules("nonexistent", None) + .expect("load unknown session"); + assert!(loaded.is_empty()); + } + + #[test] + fn a_remembered_refusal_keeps_its_reason_across_a_restart() { + use crate::session::PermissionRuleScope; + + let store = permission_store(); + let refusal = RememberedRule { + key: RuleKey { + tool_name: "shell".to_string(), + pattern: None, + }, + allow: false, + scope: PermissionRuleScope::Session, + reason: Some("this run does not allow writes".to_string()), + }; + + store + .save_rules("session-1", None, &[refusal]) + .expect("save refusal"); + + // A fresh handle on the same file, as a restarted process would open. + let reopened = SqliteRuntimeStore::new(store.path()); + let loaded = reopened + .load_rules("session-1", None) + .expect("load refusal"); + + assert_eq!(loaded.len(), 1); + assert_eq!( + loaded[0].reason.as_deref(), + Some("this run does not allow writes") + ); + } + + #[test] + fn a_rule_stored_before_the_reason_column_loads_without_one() { + use crate::session::PermissionRuleScope; + + let store = permission_store(); + // permission_rules as it stood before a refusal could carry its + // reason, written straight to the file the store is about to open. + { + let conn = Connection::open(store.path()).expect("open a pre-migration database"); + conn.execute_batch( + r#" + CREATE TABLE permission_rules ( + session_id TEXT NOT NULL, + project_id TEXT, + tool_name TEXT NOT NULL, + pattern TEXT, + allow INTEGER NOT NULL, + scope TEXT NOT NULL + ); + INSERT INTO permission_rules (session_id, project_id, tool_name, pattern, allow, scope) + VALUES ('session-1', NULL, 'shell', NULL, 0, '"session"'); + "#, + ) + .expect("write a rule in the old shape"); + } + + let loaded = store.load_rules("session-1", None).expect("load rules"); + + assert_eq!( + loaded.len(), + 1, + "opening the database must not lose the row" + ); + assert_eq!(loaded[0].key.tool_name, "shell"); + assert!(!loaded[0].allow); + assert_eq!(loaded[0].scope, PermissionRuleScope::Session); + assert_eq!( + loaded[0].reason, None, + "a rule remembered before reasons existed has none to restate" + ); + } +} diff --git a/vendor/mentra/src/runtime/task.rs b/vendor/mentra/src/runtime/task.rs new file mode 100644 index 0000000..0936569 --- /dev/null +++ b/vendor/mentra/src/runtime/task.rs @@ -0,0 +1,283 @@ +mod graph; +mod input; +mod intrinsic; +mod render; +mod store; +#[cfg(test)] +mod tests; +mod types; + +use std::{io, path::Path}; + +use serde_json::Value; +use thiserror::Error; + +use crate::runtime::{RuntimeError, store::TaskStore}; + +pub(crate) use intrinsic::TaskIntrinsicTool; +pub(crate) const TASK_REMINDER_TEXT: &str = "Reminder: use task_create, task_claim, task_update, task_list, or task_get only for persisted project-task tracking. Do not use task tools to manage persistent teammates or team protocol flows."; + +pub(crate) use graph::has_unfinished_tasks; +pub use types::{TaskItem, TaskStatus}; + +pub(crate) fn deserialize_present_nullable_string<'de, D>( + deserializer: D, +) -> Result>, D::Error> +where + D: serde::Deserializer<'de>, +{ + as serde::Deserialize>::deserialize(deserializer).map(Some) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum TaskAccess<'a> { + Lead, + LeadClaimant(&'a str), + Teammate(&'a str), +} + +#[derive(Debug, Error)] +pub(crate) enum TaskError { + #[error("Task storage I/O failed: {0}")] + Io(#[from] io::Error), + + #[error("Task serialization failed: {0}")] + Serde(#[from] serde_json::Error), + + #[error("Task validation failed: {0}")] + Validation(String), +} + +pub(crate) fn execute_with_store( + store: &S, + tool: &TaskIntrinsicTool, + input: Value, + namespace: &Path, + access: TaskAccess<'_>, +) -> Result { + match tool { + TaskIntrinsicTool::Create => { + let parsed = input::parse_task_create_input(input)?; + mutate_store_tasks(store, namespace, move |tasks| { + let task_id = tasks.iter().map(|task| task.id).max().unwrap_or(0) + 1; + tasks.push(TaskItem { + id: task_id, + subject: parsed.subject.trim().to_string(), + description: parsed.description.clone(), + status: TaskStatus::Pending, + blocked_by: Vec::new(), + blocks: Vec::new(), + owner: parsed.owner.clone(), + working_directory: parsed.working_directory.clone(), + }); + + for blocker_id in &parsed.blocked_by { + graph::add_dependency(tasks, *blocker_id, task_id) + .map_err(|error| error.to_string())?; + } + + render::serialize_pretty( + graph::find_task(tasks, task_id).map_err(|error| error.to_string())?, + ) + .map_err(|error| error.to_string()) + }) + } + TaskIntrinsicTool::Claim => { + let parsed = input::parse_task_claim_input(input)?; + let owner = access + .actor_name() + .filter(|value| !value.trim().is_empty()) + .ok_or_else(|| "Only named teammates can claim tasks".to_string())? + .trim() + .to_string(); + mutate_store_tasks(store, namespace, move |tasks| { + let claimed = match parsed.task_id { + Some(task_id) => { + let task = store::find_task_mut(tasks, task_id) + .map_err(|error| error.to_string())?; + store::validate_claimable(task, &owner) + .map_err(|error| error.to_string())?; + task.owner = owner.clone(); + task.clone() + } + None => { + let task = tasks + .iter_mut() + .find(|task| store::is_claimable(task)) + .ok_or_else(|| { + "No ready unowned tasks are available to claim".to_string() + })?; + task.owner = owner.clone(); + task.clone() + } + }; + + render::serialize_pretty(&claimed).map_err(|error| error.to_string()) + }) + } + TaskIntrinsicTool::Update => { + let parsed = input::parse_task_update_input(input)?; + mutate_store_tasks(store, namespace, move |tasks| { + let task_id = parsed.task_id; + let original_status = graph::find_task(tasks, task_id) + .map_err(|error| error.to_string())? + .status + .clone(); + store::validate_update_access( + graph::find_task(tasks, task_id).map_err(|error| error.to_string())?, + &parsed, + access, + ) + .map_err(|error| error.to_string())?; + + { + let task = + store::find_task_mut(tasks, task_id).map_err(|error| error.to_string())?; + if let Some(subject) = parsed.subject.clone() { + task.subject = subject.trim().to_string(); + } + if let Some(description) = parsed.description.clone() { + task.description = description; + } + if let Some(owner) = parsed.owner.clone() { + task.owner = owner; + } + if let Some(working_directory) = parsed.working_directory.clone() { + task.working_directory = working_directory; + } + } + + for blocker_id in &parsed.add_blocked_by { + graph::add_dependency(tasks, *blocker_id, task_id) + .map_err(|error| error.to_string())?; + } + for blocker_id in &parsed.remove_blocked_by { + graph::remove_dependency(tasks, *blocker_id, task_id) + .map_err(|error| error.to_string())?; + } + for dependent_id in &parsed.add_blocks { + graph::add_dependency(tasks, task_id, *dependent_id) + .map_err(|error| error.to_string())?; + } + for dependent_id in &parsed.remove_blocks { + graph::remove_dependency(tasks, task_id, *dependent_id) + .map_err(|error| error.to_string())?; + } + + let mut unblocked = Vec::new(); + let mut reblocked = Vec::new(); + if let Some(status) = parsed.status.clone() { + graph::apply_status_change( + tasks, + task_id, + original_status, + status, + &mut unblocked, + &mut reblocked, + ) + .map_err(|error| error.to_string())?; + } else { + store::validate_unblocked_status( + graph::find_task(tasks, task_id).map_err(|error| error.to_string())?, + ) + .map_err(|error| error.to_string())?; + } + + graph::sort_tasks(&mut unblocked); + graph::sort_tasks(&mut reblocked); + render::serialize_pretty(&render::TaskUpdateOutput { + task: graph::find_task(tasks, task_id) + .map_err(|error| error.to_string())? + .clone(), + unblocked, + reblocked, + }) + .map_err(|error| error.to_string()) + }) + } + TaskIntrinsicTool::Get => { + let parsed = input::parse_task_get_input(input)?; + let tasks = load_store_tasks(store, namespace)?; + render::serialize_pretty( + graph::find_task(&tasks, parsed.task_id).map_err(|error| error.to_string())?, + ) + .map_err(|error| error.to_string()) + } + TaskIntrinsicTool::List => { + input::parse_task_list_input(input)?; + let tasks = load_store_tasks(store, namespace)?; + let mut ready = Vec::new(); + let mut blocked = Vec::new(); + let mut in_progress = Vec::new(); + let mut completed = Vec::new(); + + for task in &tasks { + match task.status { + TaskStatus::Pending if task.blocked_by.is_empty() => ready.push(task.clone()), + TaskStatus::Pending => blocked.push(task.clone()), + TaskStatus::InProgress => in_progress.push(task.clone()), + TaskStatus::Completed => completed.push(task.clone()), + } + } + + render::serialize_pretty(&render::TaskListOutput { + tasks, + ready, + blocked, + in_progress, + completed, + }) + .map_err(|error| error.to_string()) + } + } +} + +impl<'a> TaskAccess<'a> { + pub(crate) fn actor_name(self) -> Option<&'a str> { + match self { + Self::Lead => None, + Self::LeadClaimant(name) | Self::Teammate(name) => Some(name), + } + } +} + +fn load_store_tasks( + store: &S, + namespace: &Path, +) -> Result, String> { + store + .load_tasks(namespace) + .map_err(|error| format!("Task storage failed: {error}")) +} + +fn mutate_store_tasks( + store: &S, + namespace: &Path, + mut operation: impl FnMut(&mut Vec) -> Result, +) -> Result { + let mut output = None; + let mut task_error = None; + let result = { + let mut mutation = |tasks: &mut Vec| match operation(tasks) { + Ok(value) => { + output = Some(value); + Ok(()) + } + Err(error) => { + task_error = Some(error.clone()); + Err(RuntimeError::InvalidTask(error)) + } + }; + store.mutate(namespace, &mut mutation) + }; + + match result { + Ok(()) => output.ok_or_else(|| "Task mutation did not produce a result".to_string()), + Err(_) if task_error.is_some() => Err(task_error.expect("checked above")), + Err(error) => Err(store_error(error)), + } +} + +fn store_error(error: crate::runtime::RuntimeError) -> String { + format!("Task storage failed: {error}") +} diff --git a/vendor/mentra/src/runtime/task/graph.rs b/vendor/mentra/src/runtime/task/graph.rs new file mode 100644 index 0000000..e6141e2 --- /dev/null +++ b/vendor/mentra/src/runtime/task/graph.rs @@ -0,0 +1,211 @@ +use std::collections::{BTreeSet, HashMap, HashSet}; + +use super::{ + TaskError, + types::{TaskItem, TaskStatus}, +}; + +pub(crate) fn has_unfinished_tasks(tasks: &[TaskItem]) -> bool { + tasks + .iter() + .any(|task| task.status != TaskStatus::Completed) +} + +pub(super) fn apply_status_change( + tasks: &mut [TaskItem], + task_id: u64, + original_status: TaskStatus, + next_status: TaskStatus, + unblocked: &mut Vec, + reblocked: &mut Vec, +) -> Result<(), TaskError> { + if original_status == next_status { + validate_unblocked_status(find_task(tasks, task_id)?)?; + return Ok(()); + } + + match next_status { + TaskStatus::Completed => { + { + let task = find_task_mut(tasks, task_id)?; + validate_unblocked_status(task)?; + task.status = TaskStatus::Completed; + } + + let dependents = find_task(tasks, task_id)?.blocks.clone(); + for dependent_id in dependents { + let dependent = find_task_mut(tasks, dependent_id)?; + if dependent.status == TaskStatus::Completed { + continue; + } + + let had_blocker = remove_id(&mut dependent.blocked_by, task_id); + if had_blocker && dependent.blocked_by.is_empty() { + unblocked.push(dependent.clone()); + } + } + } + TaskStatus::Pending | TaskStatus::InProgress => { + { + let task = find_task_mut(tasks, task_id)?; + task.status = next_status.clone(); + validate_unblocked_status(task)?; + } + + if original_status == TaskStatus::Completed { + let dependents = find_task(tasks, task_id)?.blocks.clone(); + for dependent_id in dependents { + let dependent = find_task_mut(tasks, dependent_id)?; + if dependent.status == TaskStatus::Completed { + continue; + } + + if insert_id(&mut dependent.blocked_by, task_id) { + reblocked.push(dependent.clone()); + } + } + } + } + } + + Ok(()) +} + +pub(super) fn add_dependency( + tasks: &mut [TaskItem], + blocker_id: u64, + dependent_id: u64, +) -> Result<(), TaskError> { + if blocker_id == dependent_id { + return Err(TaskError::Validation( + "Tasks cannot depend on themselves".to_string(), + )); + } + + let blocker_status = find_task(tasks, blocker_id)?.status.clone(); + let dependent_status = find_task(tasks, dependent_id)?.status.clone(); + + let edge_exists = find_task(tasks, blocker_id)?.blocks.contains(&dependent_id); + if !edge_exists && path_exists(tasks, dependent_id, blocker_id) { + return Err(TaskError::Validation(format!( + "Adding dependency {blocker_id} -> {dependent_id} would create a cycle" + ))); + } + + insert_id(&mut find_task_mut(tasks, blocker_id)?.blocks, dependent_id); + + if blocker_status != TaskStatus::Completed { + if dependent_status != TaskStatus::Pending { + return Err(TaskError::Validation(format!( + "Task {dependent_id} cannot have unresolved blockers while {dependent_status:?}" + ))); + } + + insert_id( + &mut find_task_mut(tasks, dependent_id)?.blocked_by, + blocker_id, + ); + } + + Ok(()) +} + +pub(super) fn remove_dependency( + tasks: &mut [TaskItem], + blocker_id: u64, + dependent_id: u64, +) -> Result<(), TaskError> { + find_task(tasks, blocker_id)?; + find_task(tasks, dependent_id)?; + + remove_id(&mut find_task_mut(tasks, blocker_id)?.blocks, dependent_id); + remove_id( + &mut find_task_mut(tasks, dependent_id)?.blocked_by, + blocker_id, + ); + Ok(()) +} + +pub(super) fn find_task(tasks: &[TaskItem], task_id: u64) -> Result<&TaskItem, TaskError> { + tasks + .iter() + .find(|task| task.id == task_id) + .ok_or_else(|| TaskError::Validation(format!("Task {task_id} does not exist"))) +} + +pub(super) fn sort_and_dedup_ids(ids: &mut Vec) { + let unique = ids.iter().copied().collect::>(); + ids.clear(); + ids.extend(unique); +} + +pub(super) fn sort_tasks(tasks: &mut [TaskItem]) { + tasks.sort_by_key(|task| task.id); +} + +fn validate_unblocked_status(task: &TaskItem) -> Result<(), TaskError> { + if matches!(task.status, TaskStatus::InProgress | TaskStatus::Completed) + && !task.blocked_by.is_empty() + { + return Err(TaskError::Validation(format!( + "Task {} cannot be {} while blocked by {:?}", + task.id, task.status, task.blocked_by + ))); + } + + Ok(()) +} + +fn find_task_mut(tasks: &mut [TaskItem], task_id: u64) -> Result<&mut TaskItem, TaskError> { + tasks + .iter_mut() + .find(|task| task.id == task_id) + .ok_or_else(|| TaskError::Validation(format!("Task {task_id} does not exist"))) +} + +fn insert_id(ids: &mut Vec, id: u64) -> bool { + if ids.contains(&id) { + return false; + } + + ids.push(id); + sort_and_dedup_ids(ids); + true +} + +fn remove_id(ids: &mut Vec, id: u64) -> bool { + let len_before = ids.len(); + ids.retain(|current| *current != id); + len_before != ids.len() +} + +fn path_exists(tasks: &[TaskItem], start: u64, goal: u64) -> bool { + if start == goal { + return true; + } + + let tasks_by_id = tasks + .iter() + .map(|task| (task.id, task)) + .collect::>(); + let mut visited = HashSet::new(); + let mut stack = vec![start]; + + while let Some(task_id) = stack.pop() { + if !visited.insert(task_id) { + continue; + } + + let Some(task) = tasks_by_id.get(&task_id) else { + continue; + }; + for next in &task.blocks { + if *next == goal { + return true; + } + stack.push(*next); + } + } + + false +} diff --git a/vendor/mentra/src/runtime/task/input.rs b/vendor/mentra/src/runtime/task/input.rs new file mode 100644 index 0000000..6d9618b --- /dev/null +++ b/vendor/mentra/src/runtime/task/input.rs @@ -0,0 +1,113 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +use super::types::TaskStatus; + +#[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub(crate) struct TaskCreateInput { + pub(crate) subject: String, + #[serde(default)] + pub(crate) description: String, + #[serde(default)] + pub(crate) owner: String, + #[serde(default)] + pub(crate) working_directory: Option, + #[serde(default)] + pub(crate) blocked_by: Vec, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub(crate) struct TaskClaimInput { + #[serde(default)] + pub(crate) task_id: Option, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub(crate) struct TaskUpdateInput { + pub(crate) task_id: u64, + #[serde(default)] + pub(crate) subject: Option, + #[serde(default)] + pub(crate) description: Option, + #[serde(default)] + pub(crate) owner: Option, + #[serde( + default, + deserialize_with = "super::deserialize_present_nullable_string" + )] + pub(crate) working_directory: Option>, + #[serde(default)] + pub(crate) status: Option, + #[serde(default)] + pub(crate) add_blocked_by: Vec, + #[serde(default)] + pub(crate) remove_blocked_by: Vec, + #[serde(default)] + pub(crate) add_blocks: Vec, + #[serde(default)] + pub(crate) remove_blocks: Vec, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub(crate) struct TaskGetInput { + pub(crate) task_id: u64, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub(crate) struct TaskListInput {} + +pub(crate) fn parse_task_create_input(input: Value) -> Result { + let parsed = serde_json::from_value::(input) + .map_err(|error| format!("Invalid task_create input: {error}"))?; + + if parsed.subject.trim().is_empty() { + return Err("Task subject must not be empty".to_string()); + } + + Ok(TaskCreateInput { + working_directory: normalize_optional_path(parsed.working_directory), + ..parsed + }) +} + +pub(crate) fn parse_task_update_input(input: Value) -> Result { + let parsed = serde_json::from_value::(input) + .map_err(|error| format!("Invalid task_update input: {error}"))?; + + if matches!(parsed.subject.as_deref(), Some(subject) if subject.trim().is_empty()) { + return Err("Task subject must not be empty".to_string()); + } + + Ok(TaskUpdateInput { + working_directory: parsed.working_directory.map(normalize_optional_path), + ..parsed + }) +} + +pub(crate) fn parse_task_claim_input(input: Value) -> Result { + serde_json::from_value::(input) + .map_err(|error| format!("Invalid task_claim input: {error}")) +} + +pub(crate) fn parse_task_get_input(input: Value) -> Result { + serde_json::from_value::(input) + .map_err(|error| format!("Invalid task_get input: {error}")) +} + +pub(crate) fn parse_task_list_input(input: Value) -> Result<(), String> { + serde_json::from_value::(input) + .map(|_| ()) + .map_err(|error| format!("Invalid task_list input: {error}")) +} + +fn normalize_optional_path(value: Option) -> Option { + value.and_then(|value| { + let trimmed = value.trim(); + (!trimmed.is_empty()).then(|| trimmed.to_string()) + }) +} diff --git a/vendor/mentra/src/runtime/task/intrinsic.rs b/vendor/mentra/src/runtime/task/intrinsic.rs new file mode 100644 index 0000000..b6a929f --- /dev/null +++ b/vendor/mentra/src/runtime/task/intrinsic.rs @@ -0,0 +1,33 @@ +#[path = "intrinsic/descriptor.rs"] +mod descriptor; +#[path = "intrinsic/execute.rs"] +mod execute; + +use async_trait::async_trait; +use strum::{Display, VariantArray}; + +use crate::tool::{RuntimeToolDescriptor, ToolContext, ToolDefinition, ToolExecutor, ToolResult}; + +#[derive(Clone, Copy, Display, VariantArray)] +#[strum(prefix = "task_")] +#[strum(serialize_all = "snake_case")] +pub enum TaskIntrinsicTool { + Create, + Claim, + Update, + List, + Get, +} + +impl ToolDefinition for TaskIntrinsicTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + descriptor::task_intrinsic_descriptor(*self) + } +} + +#[async_trait] +impl ToolExecutor for TaskIntrinsicTool { + async fn execute_mut(&self, ctx: ToolContext<'_>, input: serde_json::Value) -> ToolResult { + execute::execute_mut(*self, ctx, input) + } +} diff --git a/vendor/mentra/src/runtime/task/intrinsic/descriptor.rs b/vendor/mentra/src/runtime/task/intrinsic/descriptor.rs new file mode 100644 index 0000000..dfeac13 --- /dev/null +++ b/vendor/mentra/src/runtime/task/intrinsic/descriptor.rs @@ -0,0 +1,140 @@ +use serde_json::json; + +use crate::tool::{ + RuntimeToolDescriptor, ToolApprovalCategory, ToolCapability, ToolDurability, + ToolExecutionCategory, ToolSideEffectLevel, + internal::{RuntimeDescriptorParts, build_runtime_descriptor}, +}; + +use super::TaskIntrinsicTool; + +pub(super) fn task_intrinsic_descriptor(tool: TaskIntrinsicTool) -> RuntimeToolDescriptor { + let description = match tool { + TaskIntrinsicTool::Create => { + "Lead-oriented project planning tool. Create a persisted task." + } + TaskIntrinsicTool::Claim => { + "Claim a ready unowned persisted task for the current teammate." + } + TaskIntrinsicTool::Update => { + "Lead-oriented project planning tool. Update a persisted task and its dependency edges." + } + TaskIntrinsicTool::List => "List persisted tasks grouped by readiness.", + TaskIntrinsicTool::Get => "Get one persisted task by ID.", + }; + + let input_schema = match tool { + TaskIntrinsicTool::Create => json!({ + "type": "object", + "properties": { + "subject": { + "type": "string", + "description": "Short title for the task" + }, + "description": { + "type": "string", + "description": "Optional extra detail for the task" + }, + "owner": { + "type": "string", + "description": "Optional owner label for the task" + }, + "workingDirectory": { + "type": ["string", "null"], + "description": "Optional working directory hint for shell-based work" + }, + "blockedBy": { + "type": "array", + "items": { "type": "integer" }, + "description": "Task IDs that must finish before this task is ready" + } + }, + "required": ["subject"] + }), + TaskIntrinsicTool::Claim => json!({ + "type": "object", + "properties": { + "taskId": { + "type": "integer", + "description": "Optional explicit task identifier to claim" + } + } + }), + TaskIntrinsicTool::Update => json!({ + "type": "object", + "properties": { + "taskId": { + "type": "integer", + "description": "Stable identifier for the task" + }, + "subject": { + "type": "string", + "description": "Updated task subject" + }, + "description": { + "type": "string", + "description": "Updated task description" + }, + "owner": { + "type": "string", + "description": "Updated task owner" + }, + "workingDirectory": { + "type": ["string", "null"], + "description": "Updated working directory hint for shell-based work; pass null to clear it" + }, + "status": { + "type": "string", + "enum": ["pending", "in_progress", "completed"], + "description": "Updated task status" + }, + "addBlockedBy": { + "type": "array", + "items": { "type": "integer" }, + "description": "Add dependency edges from blocker tasks into this task" + }, + "removeBlockedBy": { + "type": "array", + "items": { "type": "integer" }, + "description": "Remove dependency edges from blocker tasks into this task" + }, + "addBlocks": { + "type": "array", + "items": { "type": "integer" }, + "description": "Add dependency edges from this task into dependent tasks" + }, + "removeBlocks": { + "type": "array", + "items": { "type": "integer" }, + "description": "Remove dependency edges from this task into dependent tasks" + } + }, + "required": ["taskId"] + }), + TaskIntrinsicTool::List => json!({ + "type": "object", + "properties": {} + }), + TaskIntrinsicTool::Get => json!({ + "type": "object", + "properties": { + "taskId": { + "type": "integer", + "description": "Stable identifier for the task" + } + }, + "required": ["taskId"] + }), + }; + + build_runtime_descriptor(RuntimeDescriptorParts { + name: tool.to_string(), + description: description.to_string(), + input_schema, + capabilities: vec![ToolCapability::TaskMutation], + side_effect_level: ToolSideEffectLevel::LocalState, + durability: ToolDurability::Persistent, + execution_category: ToolExecutionCategory::ExclusivePersistentMutation, + approval_category: ToolApprovalCategory::Default, + }) +} diff --git a/vendor/mentra/src/runtime/task/intrinsic/execute.rs b/vendor/mentra/src/runtime/task/intrinsic/execute.rs new file mode 100644 index 0000000..b42e8e8 --- /dev/null +++ b/vendor/mentra/src/runtime/task/intrinsic/execute.rs @@ -0,0 +1,52 @@ +use crate::{ + ContentBlock, + runtime::Agent, + tool::{ToolCall, ToolContext, ToolResult, internal::content_block_to_tool_result}, +}; +use strum::VariantArray; + +use super::{TaskIntrinsicTool, descriptor::task_intrinsic_descriptor}; + +pub(super) fn execute_mut( + tool: TaskIntrinsicTool, + ctx: ToolContext<'_>, + input: serde_json::Value, +) -> ToolResult { + let call = ToolCall { + id: ctx.tool_call_id.clone(), + name: task_intrinsic_descriptor(tool).provider.name, + input, + }; + let Some(result) = execute_intrinsic(ctx.agent, call) else { + return Err("Task intrinsic is not available".to_string()); + }; + content_block_to_tool_result("Task intrinsic", result) +} + +pub(super) fn execute_intrinsic(agent: &mut Agent, call: ToolCall) -> Option { + let tool = TaskIntrinsicTool::VARIANTS + .iter() + .find(|tool| task_intrinsic_descriptor(**tool).provider.name == call.name)?; + + let output = agent.execute_task_mutation(tool, call.input); + + Some(match output { + Ok(content) => match agent.refresh_tasks_from_disk() { + Ok(()) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: content.into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Task refresh failed: {error}").into(), + is_error: true, + }, + }, + Err(content) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: content.into(), + is_error: true, + }, + }) +} diff --git a/vendor/mentra/src/runtime/task/render.rs b/vendor/mentra/src/runtime/task/render.rs new file mode 100644 index 0000000..b9e436d --- /dev/null +++ b/vendor/mentra/src/runtime/task/render.rs @@ -0,0 +1,29 @@ +use serde::Serialize; + +use super::TaskError; +use super::types::TaskItem; + +#[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] +pub(super) struct TaskUpdateOutput { + pub(super) task: TaskItem, + pub(super) unblocked: Vec, + pub(super) reblocked: Vec, +} + +#[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] +pub(super) struct TaskListOutput { + pub(super) tasks: Vec, + pub(super) ready: Vec, + pub(super) blocked: Vec, + pub(super) in_progress: Vec, + pub(super) completed: Vec, +} + +pub(super) fn serialize_pretty(value: &T) -> Result +where + T: Serialize, +{ + serde_json::to_string_pretty(value).map_err(TaskError::Serde) +} diff --git a/vendor/mentra/src/runtime/task/store.rs b/vendor/mentra/src/runtime/task/store.rs new file mode 100644 index 0000000..80f1add --- /dev/null +++ b/vendor/mentra/src/runtime/task/store.rs @@ -0,0 +1,93 @@ +use super::{ + TaskAccess, TaskError, + input::TaskUpdateInput, + types::{TaskItem, TaskStatus}, +}; + +pub(super) fn validate_unblocked_status(task: &TaskItem) -> Result<(), TaskError> { + if matches!(task.status, TaskStatus::InProgress | TaskStatus::Completed) + && !task.blocked_by.is_empty() + { + return Err(TaskError::Validation(format!( + "Task {} cannot be {} while blocked by {:?}", + task.id, task.status, task.blocked_by + ))); + } + + Ok(()) +} + +pub(super) fn validate_claimable(task: &TaskItem, owner: &str) -> Result<(), TaskError> { + if !task.owner.is_empty() { + return Err(TaskError::Validation(format!( + "Task {} is already owned by '{}'", + task.id, task.owner + ))); + } + if !task.blocked_by.is_empty() { + return Err(TaskError::Validation(format!( + "Task {} is blocked by {:?} and cannot be claimed", + task.id, task.blocked_by + ))); + } + if task.status != TaskStatus::Pending { + return Err(TaskError::Validation(format!( + "Task {} is {} and cannot be claimed by '{}'", + task.id, task.status, owner + ))); + } + + Ok(()) +} + +pub(super) fn validate_update_access( + task: &TaskItem, + input: &TaskUpdateInput, + access: TaskAccess<'_>, +) -> Result<(), TaskError> { + match access { + TaskAccess::Lead | TaskAccess::LeadClaimant(_) => Ok(()), + TaskAccess::Teammate(name) if task.owner == name => { + if updates_dependencies(input) { + return Err(TaskError::Validation(format!( + "Teammate '{name}' cannot edit dependencies for task {}", + task.id + ))); + } + if let Some(owner) = &input.owner + && owner != name + { + return Err(TaskError::Validation(format!( + "Teammate '{name}' cannot reassign task {} to '{}'", + task.id, owner + ))); + } + Ok(()) + } + TaskAccess::Teammate(name) => Err(TaskError::Validation(format!( + "Teammate '{name}' cannot update task {} owned by '{}'", + task.id, task.owner + ))), + } +} + +pub(super) fn is_claimable(task: &TaskItem) -> bool { + task.status == TaskStatus::Pending && task.blocked_by.is_empty() && task.owner.is_empty() +} + +pub(super) fn find_task_mut( + tasks: &mut [TaskItem], + task_id: u64, +) -> Result<&mut TaskItem, TaskError> { + tasks + .iter_mut() + .find(|task| task.id == task_id) + .ok_or_else(|| TaskError::Validation(format!("Task {task_id} does not exist"))) +} + +fn updates_dependencies(input: &TaskUpdateInput) -> bool { + !input.add_blocked_by.is_empty() + || !input.remove_blocked_by.is_empty() + || !input.add_blocks.is_empty() + || !input.remove_blocks.is_empty() +} diff --git a/vendor/mentra/src/runtime/task/tests.rs b/vendor/mentra/src/runtime/task/tests.rs new file mode 100644 index 0000000..a69c4f7 --- /dev/null +++ b/vendor/mentra/src/runtime/task/tests.rs @@ -0,0 +1,601 @@ +use std::{ + fs, + path::PathBuf, + sync::{ + Arc, Barrier, Mutex, + atomic::{AtomicU64, AtomicUsize, Ordering}, + }, + thread, + time::{SystemTime, UNIX_EPOCH}, +}; + +use serde_json::json; + +use crate::runtime::{ + HybridRuntimeStore, RuntimeError, SqliteRuntimeStore, TaskIntrinsicTool, TaskStateSnapshot, + TaskStore, VolatileRuntimeStore, +}; + +use super::{ + TaskAccess, TaskItem, + input::{ + TaskCreateInput, TaskUpdateInput, parse_task_create_input, parse_task_list_input, + parse_task_update_input, + }, + types::TaskStatus, +}; + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +#[test] +fn concurrent_intrinsic_creates_serialize_through_task_store_mutate() { + assert_concurrent_creates(SqliteRuntimeStore::new( + temp_path("mentra-task-concurrent-sqlite".to_string()).with_extension("sqlite"), + )); + assert_concurrent_creates(VolatileRuntimeStore::new()); + assert_concurrent_creates(HybridRuntimeStore::new( + temp_path("mentra-task-concurrent-hybrid".to_string()).with_extension("sqlite"), + )); +} + +#[test] +fn custom_store_uses_the_source_compatible_default_mutate_fallback() { + let store = DefaultMutateStore::default(); + let namespace = temp_namespace("default-mutate-fallback"); + + super::execute_with_store( + &store, + &TaskIntrinsicTool::Create, + json!({ "subject": "created through default mutate" }), + &namespace, + TaskAccess::Lead, + ) + .expect("default mutate fallback persists the intrinsic mutation"); + + assert_eq!(store.load_tasks(&namespace).expect("load tasks").len(), 1); + assert_eq!(store.replacements.load(Ordering::SeqCst), 1); + + let mut rejected = |tasks: &mut Vec| { + tasks.clear(); + Err(RuntimeError::InvalidTask("reject mutation".to_string())) + }; + store + .mutate(&namespace, &mut rejected) + .expect_err("failed mutation must not replace the stored tasks"); + + assert_eq!(store.load_tasks(&namespace).expect("load tasks").len(), 1); + assert_eq!(store.replacements.load(Ordering::SeqCst), 1); +} + +#[derive(Clone, Default)] +struct DefaultMutateStore { + tasks: Arc>>, + replacements: Arc, +} + +impl TaskStore for DefaultMutateStore { + fn load_tasks(&self, _namespace: &std::path::Path) -> Result, RuntimeError> { + Ok(self.tasks.lock().expect("task store poisoned").clone()) + } + + fn capture_tasks( + &self, + namespace: &std::path::Path, + ) -> Result { + Ok(TaskStateSnapshot { + tasks: self.load_tasks(namespace)?, + }) + } + + fn restore_tasks( + &self, + namespace: &std::path::Path, + snapshot: &TaskStateSnapshot, + ) -> Result<(), RuntimeError> { + self.replace_tasks(namespace, &snapshot.tasks) + } + + fn replace_tasks( + &self, + _namespace: &std::path::Path, + tasks: &[TaskItem], + ) -> Result<(), RuntimeError> { + *self.tasks.lock().expect("task store poisoned") = tasks.to_vec(); + self.replacements.fetch_add(1, Ordering::SeqCst); + Ok(()) + } +} + +fn assert_concurrent_creates(store: impl TaskStore + Clone + 'static) { + const WRITERS: usize = 8; + + let namespace = temp_namespace("concurrent-writers"); + let barrier = Arc::new(Barrier::new(WRITERS)); + let mut writers = Vec::new(); + + // Initialize lazy stores before the simultaneous writes so this test + // isolates task-mutation serialization from schema initialization. + store.load_tasks(&namespace).expect("initialize store"); + + for index in 0..WRITERS { + let store = store.clone(); + let namespace = namespace.clone(); + let barrier = Arc::clone(&barrier); + writers.push(thread::spawn(move || { + barrier.wait(); + super::execute_with_store( + &store, + &TaskIntrinsicTool::Create, + json!({ "subject": format!("writer-{index}") }), + &namespace, + TaskAccess::Lead, + ) + .expect("create task concurrently"); + })); + } + + for writer in writers { + writer.join().expect("writer thread"); + } + + let tasks = store.load_tasks(&namespace).expect("load final tasks"); + assert_eq!(tasks.len(), WRITERS); + assert_eq!( + tasks.iter().map(|task| task.id).collect::>(), + (1..=WRITERS as u64).collect::>() + ); +} + +#[test] +fn create_and_list_group_ready_blocked_and_completed_tasks() { + let store = TaskHarness::new("grouping"); + + store.create(TaskCreateInput { + subject: "Plan".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.create(TaskCreateInput { + subject: "Build".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: vec![1], + }); + store.create(TaskCreateInput { + subject: "Review".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.update( + parse_task_update_input(json!({ + "taskId": 3, + "status": "in_progress" + })) + .expect("parse update"), + TaskAccess::Lead, + ); + store.update( + parse_task_update_input(json!({ + "taskId": 1, + "status": "completed" + })) + .expect("parse update"), + TaskAccess::Lead, + ); + + let listed = serde_json::from_str::(&store.list()).expect("parse output"); + assert_eq!(listed["ready"].as_array().expect("ready").len(), 1); + assert_eq!(listed["blocked"].as_array().expect("blocked").len(), 0); + assert_eq!( + listed["inProgress"].as_array().expect("in progress").len(), + 1 + ); + assert_eq!(listed["completed"].as_array().expect("completed").len(), 1); +} + +#[test] +fn completion_unblocks_and_reopen_reblocks_dependents() { + let store = TaskHarness::new("reblock"); + + store.create(TaskCreateInput { + subject: "A".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.create(TaskCreateInput { + subject: "B".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: vec![1], + }); + + let completed = serde_json::from_str::( + &store.update( + parse_task_update_input(json!({ + "taskId": 1, + "status": "completed" + })) + .expect("parse update"), + TaskAccess::Lead, + ), + ) + .expect("parse completed"); + assert_eq!( + completed["unblocked"].as_array().expect("unblocked").len(), + 1 + ); + + let reopened = serde_json::from_str::( + &store.update( + parse_task_update_input(json!({ + "taskId": 1, + "status": "pending" + })) + .expect("parse update"), + TaskAccess::Lead, + ), + ) + .expect("parse reopened"); + assert_eq!( + reopened["reblocked"].as_array().expect("reblocked").len(), + 1 + ); +} + +#[test] +fn adding_cycle_is_rejected() { + let store = TaskHarness::new("cycle"); + + store.create(TaskCreateInput { + subject: "A".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.create(TaskCreateInput { + subject: "B".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: vec![1], + }); + + let error = store + .try_update( + parse_task_update_input(json!({ + "taskId": 1, + "addBlockedBy": [2] + })) + .expect("parse update"), + TaskAccess::Lead, + ) + .expect_err("cycle should fail"); + assert!(error.contains("would create a cycle")); +} + +#[test] +fn blocked_task_cannot_start_or_complete() { + let store = TaskHarness::new("blocked-status"); + + store.create(TaskCreateInput { + subject: "A".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.create(TaskCreateInput { + subject: "B".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: vec![1], + }); + + let error = store + .try_update( + parse_task_update_input(json!({ + "taskId": 2, + "status": "in_progress" + })) + .expect("parse update"), + TaskAccess::Lead, + ) + .expect_err("blocked task should fail"); + assert!(error.contains("cannot be in_progress while blocked")); +} + +#[test] +fn parse_helpers_reject_bad_input() { + assert!(parse_task_create_input(json!({ "subject": "" })).is_err()); + assert!(parse_task_update_input(json!({ "taskId": 1, "bogus": true })).is_err()); + assert!(parse_task_list_input(json!({ "bogus": true })).is_err()); +} + +#[test] +fn completed_blocker_stays_out_of_unresolved_blocked_by() { + let store = TaskHarness::new("completed-blocker"); + + store.create(TaskCreateInput { + subject: "A".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.update( + parse_task_update_input(json!({ + "taskId": 1, + "status": "completed" + })) + .expect("parse update"), + TaskAccess::Lead, + ); + store.create(TaskCreateInput { + subject: "B".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: vec![1], + }); + + let tasks = store.load_all(); + assert_eq!(tasks[1].status, TaskStatus::Pending); + assert!(tasks[1].blocked_by.is_empty()); + assert_eq!(tasks[0].blocks, vec![2]); +} + +#[test] +fn claim_first_ready_unowned_task() { + let store = TaskHarness::new("claim-first"); + store.create(TaskCreateInput { + subject: "A".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.create(TaskCreateInput { + subject: "B".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: vec![1], + }); + + let claimed = serde_json::from_str::(&store.claim(None, "alice")) + .expect("parse claimed"); + assert_eq!(claimed["id"].as_u64(), Some(1)); + assert_eq!(claimed["owner"].as_str(), Some("alice")); +} + +#[test] +fn claim_explicit_task_id() { + let store = TaskHarness::new("claim-explicit"); + store.create(TaskCreateInput { + subject: "A".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.create(TaskCreateInput { + subject: "B".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + + let claimed = serde_json::from_str::(&store.claim(Some(2), "bob")) + .expect("parse claimed"); + assert_eq!(claimed["id"].as_u64(), Some(2)); + assert_eq!(claimed["owner"].as_str(), Some("bob")); +} + +#[test] +fn claim_rejects_unclaimable_tasks() { + let store = TaskHarness::new("claim-reject"); + store.create(TaskCreateInput { + subject: "A".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.create(TaskCreateInput { + subject: "B".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: vec![1], + }); + + let blocked = store.try_claim(Some(2), "alice").expect_err("blocked task"); + assert!(blocked.contains("cannot be claimed")); + + store.claim(Some(1), "alice"); + let owned = store.try_claim(Some(1), "bob").expect_err("owned task"); + assert!(owned.contains("already owned")); + + let missing = store + .try_claim(Some(99), "alice") + .expect_err("missing task"); + assert!(missing.contains("does not exist")); + + let store = TaskHarness::new("claim-status"); + store.create(TaskCreateInput { + subject: "C".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.update( + parse_task_update_input(json!({ + "taskId": 1, + "status": "in_progress" + })) + .expect("parse update"), + TaskAccess::Lead, + ); + let in_progress = store + .try_claim(Some(1), "alice") + .expect_err("in progress task"); + assert!(in_progress.contains("cannot be claimed")); + + let store = TaskHarness::new("claim-completed"); + store.create(TaskCreateInput { + subject: "D".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.update( + parse_task_update_input(json!({ + "taskId": 1, + "status": "completed" + })) + .expect("parse update"), + TaskAccess::Lead, + ); + let completed = store + .try_claim(Some(1), "alice") + .expect_err("completed task"); + assert!(completed.contains("cannot be claimed")); +} + +#[test] +fn teammate_cannot_edit_task_dependencies() { + let store = TaskHarness::new("teammate-deps"); + store.create(TaskCreateInput { + subject: "Owned".to_string(), + description: String::new(), + owner: "alice".to_string(), + working_directory: None, + blocked_by: Vec::new(), + }); + store.create(TaskCreateInput { + subject: "Other".to_string(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + }); + + let error = store + .try_update( + parse_task_update_input(json!({ + "taskId": 1, + "addBlocks": [2] + })) + .expect("parse update"), + TaskAccess::Teammate("alice"), + ) + .expect_err("dependency edit should fail"); + assert!(error.contains("cannot edit dependencies")); +} + +struct TaskHarness { + store: SqliteRuntimeStore, + namespace: PathBuf, +} + +impl TaskHarness { + fn new(label: &str) -> Self { + Self { + store: temp_store(label), + namespace: temp_namespace(label), + } + } + + fn create(&self, input: TaskCreateInput) -> String { + self.try_create(input).expect("create task") + } + + fn try_create(&self, input: TaskCreateInput) -> Result { + super::execute_with_store( + &self.store, + &super::TaskIntrinsicTool::Create, + serde_json::to_value(input).expect("serialize task create input"), + self.namespace.as_path(), + TaskAccess::Lead, + ) + } + + fn update(&self, input: TaskUpdateInput, access: TaskAccess<'_>) -> String { + self.try_update(input, access).expect("update task") + } + + fn try_update(&self, input: TaskUpdateInput, access: TaskAccess<'_>) -> Result { + super::execute_with_store( + &self.store, + &super::TaskIntrinsicTool::Update, + serde_json::to_value(input).expect("serialize task update input"), + self.namespace.as_path(), + access, + ) + } + + fn claim(&self, task_id: Option, owner: &str) -> String { + self.try_claim(task_id, owner).expect("claim task") + } + + fn try_claim(&self, task_id: Option, owner: &str) -> Result { + super::execute_with_store( + &self.store, + &TaskIntrinsicTool::Claim, + json!({ "taskId": task_id }), + self.namespace.as_path(), + TaskAccess::Teammate(owner), + ) + } + + fn list(&self) -> String { + super::execute_with_store( + &self.store, + &TaskIntrinsicTool::List, + json!({}), + self.namespace.as_path(), + TaskAccess::Lead, + ) + .expect("list tasks") + } + + fn load_all(&self) -> Vec { + self.store + .load_tasks(self.namespace.as_path()) + .expect("load tasks") + } +} + +fn temp_namespace(label: &str) -> PathBuf { + let path = temp_path(format!("mentra-task-graph-{label}")); + fs::create_dir_all(&path).expect("create temp namespace dir"); + path +} + +fn temp_store(label: &str) -> SqliteRuntimeStore { + SqliteRuntimeStore::new( + temp_path(format!("mentra-task-store-{label}")).with_extension("sqlite"), + ) +} + +fn temp_path(label: String) -> PathBuf { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + std::env::temp_dir().join(format!("{label}-{timestamp}-{unique}")) +} diff --git a/vendor/mentra/src/runtime/task/types.rs b/vendor/mentra/src/runtime/task/types.rs new file mode 100644 index 0000000..c4557e6 --- /dev/null +++ b/vendor/mentra/src/runtime/task/types.rs @@ -0,0 +1,31 @@ +use serde::{Deserialize, Serialize}; +use strum::Display; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default, Display)] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum TaskStatus { + #[default] + Pending, + InProgress, + Completed, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct TaskItem { + pub id: u64, + pub subject: String, + #[serde(default)] + pub description: String, + #[serde(default)] + pub status: TaskStatus, + #[serde(default)] + pub blocked_by: Vec, + #[serde(default)] + pub blocks: Vec, + #[serde(default)] + pub owner: String, + #[serde(default)] + pub working_directory: Option, +} diff --git a/vendor/mentra/src/runtime/task_board.rs b/vendor/mentra/src/runtime/task_board.rs new file mode 100644 index 0000000..8061cd8 --- /dev/null +++ b/vendor/mentra/src/runtime/task_board.rs @@ -0,0 +1,229 @@ +use std::path::PathBuf; + +use serde::{Deserialize, Serialize, de::DeserializeOwned}; +use thiserror::Error; + +use super::{RuntimeHandle, TaskItem, TaskStatus, task::TaskAccess}; +use crate::runtime::task::TaskIntrinsicTool; + +/// Input for creating a persisted task through [`TaskBoard`]. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct NewTask { + pub subject: String, + #[serde(default)] + pub description: String, + #[serde(default)] + pub owner: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub working_directory: Option, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub blocked_by: Vec, +} + +impl NewTask { + pub fn new(subject: impl Into) -> Self { + Self { + subject: subject.into(), + description: String::new(), + owner: String::new(), + working_directory: None, + blocked_by: Vec::new(), + } + } +} + +/// Typed fields accepted by [`TaskBoard::update`]. +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct TaskPatch { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub subject: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub description: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub owner: Option, + #[serde( + default, + skip_serializing_if = "Option::is_none", + deserialize_with = "super::task::deserialize_present_nullable_string" + )] + pub working_directory: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub status: Option, +} + +/// Error returned by typed task-board operations. +#[derive(Debug, Error)] +pub enum TaskBoardError { + #[error("{0}")] + Operation(String), + #[error("task operation returned an incompatible result: {0}")] + InvalidResult(#[source] serde_json::Error), +} + +#[derive(Debug, Clone)] +enum BoardAccess { + Lead, + Teammate(String), +} + +/// Cloneable typed façade over Mentra's dependency-aware task board. +/// +/// The façade deliberately delegates every operation to the builtin task +/// executor, preserving one implementation of access checks, DAG validation, +/// status propagation, and storage transactions. A `TaskBoard` never caches +/// task items: every read observes the store when the method is called. +#[derive(Clone)] +pub struct TaskBoard { + runtime: RuntimeHandle, + namespace: PathBuf, + access: BoardAccess, +} + +impl TaskBoard { + pub(crate) fn lead(runtime: RuntimeHandle, namespace: PathBuf) -> Self { + Self { + runtime, + namespace, + access: BoardAccess::Lead, + } + } + + pub(crate) fn agent( + runtime: RuntimeHandle, + namespace: PathBuf, + name: String, + is_teammate: bool, + ) -> Self { + Self { + runtime, + namespace, + access: if is_teammate { + BoardAccess::Teammate(name) + } else { + BoardAccess::Lead + }, + } + } + + pub fn create(&self, spec: NewTask) -> Result { + self.execute(TaskIntrinsicTool::Create, spec, self.access()) + } + + pub fn get(&self, id: u64) -> Result { + self.execute( + TaskIntrinsicTool::Get, + serde_json::json!({ "taskId": id }), + self.access(), + ) + } + + pub fn list(&self) -> Result, TaskBoardError> { + let result: TaskListResult = self.execute( + TaskIntrinsicTool::List, + serde_json::json!({}), + self.access(), + )?; + Ok(result.tasks) + } + + pub fn update(&self, id: u64, patch: TaskPatch) -> Result { + let mut input = serde_json::to_value(patch).map_err(TaskBoardError::InvalidResult)?; + input + .as_object_mut() + .expect("TaskPatch serializes as an object") + .insert("taskId".to_string(), serde_json::json!(id)); + let result: TaskUpdateResult = + self.execute(TaskIntrinsicTool::Update, input, self.access())?; + Ok(result.task) + } + + /// Claims a ready task for `owner`. + /// + /// Runtime-scoped boards have lead access and therefore require an + /// explicit claimant instead of pretending the host is a teammate. + /// Teammate-scoped boards reject any owner other than the agent itself. + pub fn claim(&self, id: Option, owner: &str) -> Result { + let owner = owner.trim(); + if owner.is_empty() { + return Err(TaskBoardError::Operation( + "Task claimant must not be empty".to_string(), + )); + } + + let access = match &self.access { + BoardAccess::Lead => TaskAccess::LeadClaimant(owner), + BoardAccess::Teammate(name) if name == owner => TaskAccess::Teammate(name), + BoardAccess::Teammate(name) => { + return Err(TaskBoardError::Operation(format!( + "Teammate '{name}' cannot claim a task for '{owner}'" + ))); + } + }; + self.execute( + TaskIntrinsicTool::Claim, + serde_json::json!({ "taskId": id }), + access, + ) + } + + pub fn add_dependency(&self, blocker: u64, dependent: u64) -> Result { + self.update_dependency(dependent, "addBlockedBy", blocker) + } + + pub fn remove_dependency( + &self, + blocker: u64, + dependent: u64, + ) -> Result { + self.update_dependency(dependent, "removeBlockedBy", blocker) + } + + fn update_dependency( + &self, + task_id: u64, + field: &str, + related_id: u64, + ) -> Result { + let result: TaskUpdateResult = self.execute( + TaskIntrinsicTool::Update, + serde_json::json!({ "taskId": task_id, (field): [related_id] }), + self.access(), + )?; + Ok(result.task) + } + + fn access(&self) -> TaskAccess<'_> { + match &self.access { + BoardAccess::Lead => TaskAccess::Lead, + BoardAccess::Teammate(name) => TaskAccess::Teammate(name), + } + } + + fn execute( + &self, + tool: TaskIntrinsicTool, + input: impl Serialize, + access: TaskAccess<'_>, + ) -> Result { + let input = serde_json::to_value(input).map_err(TaskBoardError::InvalidResult)?; + let output = self + .runtime + .execute_task_mutation(&tool, input, &self.namespace, access) + .map_err(TaskBoardError::Operation)?; + serde_json::from_str(&output).map_err(TaskBoardError::InvalidResult) + } +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct TaskListResult { + tasks: Vec, +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct TaskUpdateResult { + task: TaskItem, +} diff --git a/vendor/mentra/src/runtime/volatile_store.rs b/vendor/mentra/src/runtime/volatile_store.rs new file mode 100644 index 0000000..17e89bb --- /dev/null +++ b/vendor/mentra/src/runtime/volatile_store.rs @@ -0,0 +1,508 @@ +//! An in-memory `RuntimeStore` for genuinely ephemeral runs. +//! +//! [`VolatileRuntimeStore`] satisfies the full `RuntimeStore` composition +//! (agent records, runs, tasks, audit events, leases, permission rules, team +//! state, background-task notifications, and long-term memory) entirely in +//! process memory. It never touches disk: no SQLite file is opened, no +//! transcript `.jsonl` snapshot is written, no directory is created. +//! Dropping every clone of the store leaves nothing behind. +//! +//! ## Isolation across runs +//! +//! The **recommended pattern** is to construct a fresh store per run — +//! [`VolatileRuntimeStore::new`] is trivial (no I/O, a handful of empty +//! collections behind one `Arc>`), so building one per ask is +//! cheap. A fresh store per run isolates runs by construction: there is +//! nothing to leak because nothing is shared. +//! +//! A host that instead **retains** one `VolatileRuntimeStore` across +//! multiple sequential runs (for example inside a pooled `Runtime`) does not +//! get that isolation for free. Several `RuntimeStore` methods have no +//! per-run scope in their signature: [`AgentStore::list_agents`] lists every +//! agent record the store has ever seen, and the `TeamStore`/`TaskStore` +//! seams are keyed by `team_dir`/`tasks_dir` paths and agent *names*, which +//! `AgentConfig::default()` gives the same value across every agent built in +//! one process. A retained store therefore behaves like a shared database: +//! two runs that use the same team directory, tasks directory, or agent name +//! will see each other's records, exactly as two `Agent::run`s pointed at +//! the same `SqliteRuntimeStore` path would. +//! +//! [`VolatileRuntimeStore::reset`] is the explicit isolation seam for that +//! case: it clears all in-memory state atomically, on every clone (clones +//! share the same backing state, like [`SqliteRuntimeStore`](super::SqliteRuntimeStore)'s clones share the +//! same file). A host that retains one store across runs must call +//! `reset()` between runs to get the same no-cross-run-visibility guarantee +//! that constructing a fresh store gives automatically. + +mod background; +mod memory; +mod permission; +mod task; +mod team; + +use std::{ + collections::HashMap, + path::Path, + sync::{Arc, Mutex, MutexGuard}, + time::{Duration, Instant}, +}; + +use background::BackgroundState; +use memory::MemoryState; +use permission::PermissionState; +use task::TaskState; +use team::TeamState; + +use crate::memory::journal::AgentMemoryState; + +use super::{AgentStore, AuditStore, LeaseStore, LoadedAgentState, PersistedAgentRecord, RunStore}; +use crate::runtime::RuntimeError; + +/// Converts a filesystem path into the string key used to namespace +/// path-scoped state (`tasks_dir`, `team_dir`). Mirrors +/// [`SqliteRuntimeStore`](super::SqliteRuntimeStore)'s use of the path's +/// string form as a SQL key. +fn path_key(path: &Path) -> String { + path.to_string_lossy().into_owned() +} + +/// Delivery/notification lifecycle shared by the team inbox and the +/// background-task notification queue: an entry starts `Pending`, moves to +/// `Inflight` while a round is actively reading it, and ends `Acked` (or is +/// requeued back to `Pending` when the run that read it fails). +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum DeliveryState { + Pending, + Inflight, + Acked, +} + +struct RunRecord { + state: String, + error: Option, +} + +struct LeaseEntry { + owner: String, + expires_at: Instant, +} + +#[derive(Default)] +struct VolatileState { + agents: HashMap, + agent_memory: HashMap, + agent_order: Vec, + runs: HashMap, + next_run_id: u64, + leases: HashMap, + tasks: TaskState, + team: TeamState, + background: BackgroundState, + permissions: PermissionState, + memory: MemoryState, + #[cfg(test)] + fail_next_agent_record_save: bool, + #[cfg(test)] + recovery_preparations: usize, +} + +/// An in-memory [`RuntimeStore`](super::RuntimeStore) that leaves no durable +/// trace. See the module docs for the isolation contract when one instance +/// is retained across multiple runs. +#[derive(Clone)] +pub struct VolatileRuntimeStore { + state: Arc>, +} + +impl VolatileRuntimeStore { + /// Creates an empty volatile store. Construction is trivial (no I/O) — + /// building a fresh instance per run is the recommended pattern; see the + /// module docs for the retained-store alternative. + pub fn new() -> Self { + Self { + state: Arc::new(Mutex::new(VolatileState::default())), + } + } + + /// Clears all in-memory state on every clone of this store. + /// + /// Call this between runs when retaining one `VolatileRuntimeStore` + /// across multiple sequential runs (for example inside a pooled + /// `Runtime`) to prevent one run's records from being visible to the + /// next. See the module docs for why a retained store needs this. + pub fn reset(&self) { + *self.lock() = VolatileState::default(); + } + + #[cfg(test)] + pub(crate) fn fail_next_agent_record_save(&self) { + self.lock().fail_next_agent_record_save = true; + } + + /// How many times recovery has been prepared on this store, counted across + /// every clone. Lets a test see *when* a runtime prepares recovery, and on + /// which store. + #[cfg(test)] + pub(crate) fn recovery_preparations(&self) -> usize { + self.lock().recovery_preparations + } + + fn lock(&self) -> MutexGuard<'_, VolatileState> { + self.state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + } +} + +impl Default for VolatileRuntimeStore { + fn default() -> Self { + Self::new() + } +} + +impl AgentStore for VolatileRuntimeStore { + fn allows_disk_artifacts(&self) -> bool { + false + } + + fn prepare_recovery(&self) -> Result<(), RuntimeError> { + // Nothing to recover: a volatile store never survives a process + // restart, so there is no interrupted state to reconcile. The count + // is the only trace, and exists so a test can pin down when the + // runtime prepares recovery. + #[cfg(test)] + { + self.lock().recovery_preparations += 1; + } + Ok(()) + } + + fn create_agent( + &self, + record: &PersistedAgentRecord, + memory: &AgentMemoryState, + ) -> Result<(), RuntimeError> { + let mut state = self.lock(); + if !state.agents.contains_key(&record.id) { + state.agent_order.push(record.id.clone()); + } + state.agents.insert(record.id.clone(), record.clone()); + state.agent_memory.insert(record.id.clone(), memory.clone()); + Ok(()) + } + + fn save_agent_record(&self, record: &PersistedAgentRecord) -> Result<(), RuntimeError> { + let mut state = self.lock(); + #[cfg(test)] + if std::mem::take(&mut state.fail_next_agent_record_save) { + return Err(RuntimeError::Store( + "injected agent-record persistence failure".to_string(), + )); + } + if !state.agents.contains_key(&record.id) { + state.agent_order.push(record.id.clone()); + } + state.agents.insert(record.id.clone(), record.clone()); + Ok(()) + } + + fn save_agent_memory( + &self, + agent_id: &str, + memory: &AgentMemoryState, + ) -> Result<(), RuntimeError> { + self.lock() + .agent_memory + .insert(agent_id.to_string(), memory.clone()); + Ok(()) + } + + fn load_agent(&self, agent_id: &str) -> Result, RuntimeError> { + let state = self.lock(); + let Some(record) = state.agents.get(agent_id).cloned() else { + return Ok(None); + }; + let Some(memory) = state.agent_memory.get(agent_id).cloned() else { + return Err(RuntimeError::Store(format!( + "Agent '{agent_id}' is missing persisted memory" + ))); + }; + Ok(Some(LoadedAgentState { record, memory })) + } + + fn list_agents(&self) -> Result, RuntimeError> { + let state = self.lock(); + state + .agent_order + .iter() + .map(|id| { + let record = state + .agents + .get(id) + .cloned() + .ok_or_else(|| RuntimeError::Store(format!("Agent '{id}' disappeared")))?; + let memory = state.agent_memory.get(id).cloned().ok_or_else(|| { + RuntimeError::Store(format!("Agent '{id}' is missing persisted memory")) + })?; + Ok(LoadedAgentState { record, memory }) + }) + .collect() + } + + fn list_agents_by_runtime( + &self, + runtime_identifier: &str, + ) -> Result, RuntimeError> { + Ok(self + .list_agents()? + .into_iter() + .filter(|loaded| loaded.record.runtime_identifier == runtime_identifier) + .collect()) + } +} + +impl RunStore for VolatileRuntimeStore { + fn start_run(&self, _agent_id: &str) -> Result { + let mut state = self.lock(); + state.next_run_id += 1; + let run_id = format!("volatile-run-{}", state.next_run_id); + state.runs.insert( + run_id.clone(), + RunRecord { + state: "running".to_string(), + error: None, + }, + ); + Ok(run_id) + } + + fn update_run_state( + &self, + run_id: &str, + run_state: &str, + error: Option<&str>, + ) -> Result<(), RuntimeError> { + // Matches the default store's UPDATE-affecting-zero-rows behavior: + // updating an unknown run id is a silent no-op, not an error. + let mut state = self.lock(); + if let Some(run) = state.runs.get_mut(run_id) { + run.state = run_state.to_string(); + run.error = error.map(str::to_string); + } + Ok(()) + } + + fn finish_run(&self, run_id: &str) -> Result<(), RuntimeError> { + self.update_run_state(run_id, "finished", None) + } + + fn fail_run(&self, run_id: &str, error: &str) -> Result<(), RuntimeError> { + self.update_run_state(run_id, "failed", Some(error)) + } +} + +impl AuditStore for VolatileRuntimeStore { + fn record_audit_event( + &self, + _scope: &str, + _event_type: &str, + _payload: serde_json::Value, + ) -> Result<(), RuntimeError> { + // `AuditStore` has no reader method — nothing in mentra ever reads + // an audit event back. The volatile profile accepts the write and + // discards it rather than growing an in-memory log nobody consumes. + Ok(()) + } +} + +impl LeaseStore for VolatileRuntimeStore { + fn acquire_lease(&self, key: &str, owner: &str, ttl: Duration) -> Result { + let mut state = self.lock(); + let now = Instant::now(); + state.leases.retain(|_, lease| lease.expires_at > now); + if state.leases.contains_key(key) { + return Ok(false); + } + state.leases.insert( + key.to_string(), + LeaseEntry { + owner: owner.to_string(), + expires_at: now + ttl, + }, + ); + Ok(true) + } + + fn release_lease(&self, key: &str, owner: &str) -> Result<(), RuntimeError> { + let mut state = self.lock(); + if state + .leases + .get(key) + .is_some_and(|lease| lease.owner == owner) + { + state.leases.remove(key); + } + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use crate::{ + agent::{AgentConfig, AgentStatus}, + provider::ProviderId, + runtime::{AgentStore, LeaseStore, RunStore, RuntimeError}, + }; + + use super::{AgentMemoryState, PersistedAgentRecord, VolatileRuntimeStore}; + + fn agent_record(id: &str) -> PersistedAgentRecord { + PersistedAgentRecord { + id: id.to_string(), + runtime_identifier: "test-runtime".to_string(), + name: format!("agent-{id}"), + model: "test-model".to_string(), + provider_id: ProviderId::new("test"), + config: AgentConfig::default(), + hidden_tools: Default::default(), + max_rounds: None, + teammate_identity: None, + rounds_since_task: 0, + idle_requested: false, + status: AgentStatus::default(), + subagents: Vec::new(), + } + } + + #[test] + fn create_agent_then_load_round_trips() { + let store = VolatileRuntimeStore::new(); + let record = agent_record("agent-1"); + let memory = AgentMemoryState::default(); + + store.create_agent(&record, &memory).expect("create agent"); + + let loaded = store + .load_agent("agent-1") + .expect("load agent") + .expect("agent present"); + assert_eq!(loaded.record.id, "agent-1"); + assert_eq!(loaded.record.name, "agent-agent-1"); + } + + #[test] + fn save_agent_record_upserts_without_prior_create() { + let store = VolatileRuntimeStore::new(); + let mut record = agent_record("agent-2"); + store.save_agent_record(&record).expect("save record"); + + let err = store + .load_agent("agent-2") + .expect_err("memory should be missing until it is saved"); + assert!(matches!(err, RuntimeError::Store(_))); + + record.name = "renamed".to_string(); + store + .save_agent_record(&record) + .expect("save updated record"); + store + .save_agent_memory("agent-2", &AgentMemoryState::default()) + .expect("save memory"); + + let loaded = store + .load_agent("agent-2") + .expect("load agent") + .expect("agent present"); + assert_eq!(loaded.record.name, "renamed"); + } + + #[test] + fn list_agents_reflects_creation_order() { + let store = VolatileRuntimeStore::new(); + let memory = AgentMemoryState::default(); + store + .create_agent(&agent_record("first"), &memory) + .expect("create first"); + store + .create_agent(&agent_record("second"), &memory) + .expect("create second"); + + let ids: Vec<_> = store + .list_agents() + .expect("list agents") + .into_iter() + .map(|loaded| loaded.record.id) + .collect(); + assert_eq!(ids, vec!["first".to_string(), "second".to_string()]); + } + + #[test] + fn lease_round_trips_and_frees_on_release() { + let store = VolatileRuntimeStore::new(); + assert!( + store + .acquire_lease("agent:x", "owner-1", Duration::from_secs(60)) + .expect("acquire") + ); + assert!( + !store + .acquire_lease("agent:x", "owner-2", Duration::from_secs(60)) + .expect("second acquire"), + "lease should still be held by owner-1" + ); + + store.release_lease("agent:x", "owner-1").expect("release"); + assert!( + store + .acquire_lease("agent:x", "owner-2", Duration::from_secs(60)) + .expect("reacquire after release") + ); + } + + #[test] + fn reset_clears_all_state() { + let store = VolatileRuntimeStore::new(); + let memory = AgentMemoryState::default(); + store + .create_agent(&agent_record("agent-1"), &memory) + .expect("create agent"); + store + .acquire_lease("agent:agent-1", "owner", Duration::from_secs(60)) + .expect("acquire lease"); + let run_id = store.start_run("agent-1").expect("start run"); + + store.reset(); + + assert!(store.list_agents().expect("list agents").is_empty()); + assert!( + store + .acquire_lease("agent:agent-1", "owner-2", Duration::from_secs(60)) + .expect("lease is free after reset") + ); + // The run id from before reset() no longer resolves to anything; + // updating it is a silent no-op, matching the default store. + store + .update_run_state(&run_id, "finished", None) + .expect("update on a stale run id is a no-op"); + } + + #[test] + fn cloned_store_shares_state() { + let store = VolatileRuntimeStore::new(); + let clone = store.clone(); + let memory = AgentMemoryState::default(); + + clone + .create_agent(&agent_record("shared"), &memory) + .expect("create via clone"); + + assert!( + store + .load_agent("shared") + .expect("load via original") + .is_some() + ); + } +} diff --git a/vendor/mentra/src/runtime/volatile_store/background.rs b/vendor/mentra/src/runtime/volatile_store/background.rs new file mode 100644 index 0000000..011b757 --- /dev/null +++ b/vendor/mentra/src/runtime/volatile_store/background.rs @@ -0,0 +1,213 @@ +use std::collections::HashMap; + +use crate::{ + background::{BackgroundNotification, BackgroundStore, BackgroundTaskSummary}, + runtime::RuntimeError, +}; + +use super::{DeliveryState, VolatileRuntimeStore}; + +struct BackgroundJobEntry { + task: BackgroundTaskSummary, + notification_state: DeliveryState, +} + +/// Background-task notifications keyed by `(agent_id, task_id)`, mirroring +/// the default store's `background_jobs` table. +#[derive(Default)] +pub(super) struct BackgroundState { + jobs: HashMap<(String, String), BackgroundJobEntry>, +} + +fn notification_state_from_raw(value: i64) -> DeliveryState { + match value { + 0 => DeliveryState::Pending, + 1 => DeliveryState::Inflight, + _ => DeliveryState::Acked, + } +} + +impl BackgroundStore for VolatileRuntimeStore { + fn load_background_tasks( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + let state = self.lock(); + let mut tasks: Vec<_> = state + .background + .jobs + .iter() + .filter(|((aid, _), _)| aid == agent_id) + .map(|(_, entry)| entry.task.clone()) + .collect(); + tasks.sort_by(|a, b| a.id.cmp(&b.id)); + Ok(tasks) + } + + fn upsert_background_task( + &self, + agent_id: &str, + task: &BackgroundTaskSummary, + notification_state: i64, + ) -> Result<(), RuntimeError> { + self.lock().background.jobs.insert( + (agent_id.to_string(), task.id.clone()), + BackgroundJobEntry { + task: task.clone(), + notification_state: notification_state_from_raw(notification_state), + }, + ); + Ok(()) + } + + fn drain_background_notifications( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + let mut state = self.lock(); + let mut out = Vec::new(); + for ((aid, _id), entry) in state.background.jobs.iter_mut() { + if aid == agent_id && entry.notification_state == DeliveryState::Pending { + entry.notification_state = DeliveryState::Inflight; + out.push(BackgroundNotification { + task_id: entry.task.id.clone(), + command: entry.task.command.clone(), + cwd: entry.task.cwd.clone(), + status: entry.task.status.clone(), + output_preview: entry + .task + .output_preview + .clone() + .unwrap_or_else(|| "(no output)".to_string()), + }); + } + } + Ok(out) + } + + fn has_deliverable_background_notifications( + &self, + agent_id: &str, + ) -> Result { + Ok(self.lock().background.jobs.iter().any(|((aid, _), entry)| { + aid == agent_id && entry.notification_state == DeliveryState::Pending + })) + } + + fn has_pending_background_notifications(&self, agent_id: &str) -> Result { + Ok(self.lock().background.jobs.iter().any(|((aid, _), entry)| { + aid == agent_id + && matches!( + entry.notification_state, + DeliveryState::Pending | DeliveryState::Inflight + ) + })) + } + + fn ack_background_notifications(&self, agent_id: &str) -> Result<(), RuntimeError> { + let mut state = self.lock(); + for ((aid, _), entry) in state.background.jobs.iter_mut() { + if aid == agent_id && entry.notification_state == DeliveryState::Inflight { + entry.notification_state = DeliveryState::Acked; + } + } + Ok(()) + } + + fn requeue_background_notifications(&self, agent_id: &str) -> Result<(), RuntimeError> { + let mut state = self.lock(); + for ((aid, _), entry) in state.background.jobs.iter_mut() { + if aid == agent_id && entry.notification_state == DeliveryState::Inflight { + entry.notification_state = DeliveryState::Pending; + } + } + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use std::path::PathBuf; + + use crate::background::{BackgroundStore, BackgroundTaskStatus, BackgroundTaskSummary}; + + use super::super::VolatileRuntimeStore; + + fn summary(id: &str) -> BackgroundTaskSummary { + BackgroundTaskSummary { + id: id.to_string(), + command: "echo hi".to_string(), + cwd: PathBuf::from("/tmp"), + status: BackgroundTaskStatus::Running, + output_preview: None, + } + } + + #[test] + fn drain_then_ack_notifications() { + let store = VolatileRuntimeStore::new(); + store + .upsert_background_task("agent-1", &summary("bg-1"), 0) + .expect("seed pending task"); + + assert!( + store + .has_deliverable_background_notifications("agent-1") + .expect("has deliverable") + ); + + let drained = store + .drain_background_notifications("agent-1") + .expect("drain notifications"); + assert_eq!(drained.len(), 1); + assert_eq!(drained[0].task_id, "bg-1"); + + // Draining moves the notification to in-flight; it is no longer + // freshly deliverable, but it is still pending overall. + assert!( + !store + .has_deliverable_background_notifications("agent-1") + .expect("has deliverable after drain") + ); + assert!( + store + .has_pending_background_notifications("agent-1") + .expect("has pending after drain") + ); + + store + .ack_background_notifications("agent-1") + .expect("ack notifications"); + assert!( + !store + .has_pending_background_notifications("agent-1") + .expect("has pending after ack") + ); + } + + #[test] + fn background_tasks_are_scoped_per_agent() { + let store = VolatileRuntimeStore::new(); + store + .upsert_background_task("agent-a", &summary("bg-1"), 2) + .expect("seed agent a"); + store + .upsert_background_task("agent-b", &summary("bg-1"), 2) + .expect("seed agent b"); + + assert_eq!( + store + .load_background_tasks("agent-a") + .expect("load agent a") + .len(), + 1 + ); + assert_eq!( + store + .load_background_tasks("agent-b") + .expect("load agent b") + .len(), + 1 + ); + } +} diff --git a/vendor/mentra/src/runtime/volatile_store/memory.rs b/vendor/mentra/src/runtime/volatile_store/memory.rs new file mode 100644 index 0000000..bd0b402 --- /dev/null +++ b/vendor/mentra/src/runtime/volatile_store/memory.rs @@ -0,0 +1,340 @@ +use std::collections::HashMap; + +use crate::{ + memory::{ + MemoryCursor, MemoryListPage, MemoryListRequest, MemoryListSort, MemoryRecord, + MemorySearchRequest, MemoryStore, + }, + runtime::RuntimeError, +}; + +use super::VolatileRuntimeStore; + +/// Long-term memory records and per-agent ingest cursors, mirroring the +/// default store's `long_term_memory` / `long_term_memory_cursor` tables. +/// +/// Search here is a simple case-insensitive substring match over each query +/// token, not the default store's BM25-ranked full-text search — the +/// volatile profile favors simplicity over search quality for ephemeral +/// runs. +#[derive(Default)] +pub(super) struct MemoryState { + records: HashMap, + cursors: HashMap, +} + +fn query_tokens(query: &str) -> Vec { + query + .split(|ch: char| !ch.is_alphanumeric()) + .filter(|token| !token.is_empty()) + .map(str::to_lowercase) + .collect() +} + +impl MemoryStore for VolatileRuntimeStore { + fn upsert_records(&self, records: &[MemoryRecord]) -> Result<(), RuntimeError> { + let mut state = self.lock(); + for record in records { + state + .memory + .records + .insert(record.record_id.clone(), record.clone()); + } + Ok(()) + } + + fn search_records_with_options( + &self, + request: &MemorySearchRequest, + ) -> Result, RuntimeError> { + if request.limit == 0 { + return Ok(Vec::new()); + } + let tokens = query_tokens(&request.query); + if tokens.is_empty() { + return Ok(Vec::new()); + } + + let state = self.lock(); + let mut matches: Vec = state + .memory + .records + .values() + .filter(|record| { + if record.agent_id != request.agent_id { + return false; + } + if request.filter.kind.is_some_and(|kind| record.kind != kind) + || request + .filter + .pinned + .is_some_and(|pinned| record.pinned != pinned) + || request + .filter + .source + .as_ref() + .is_some_and(|source| record.source.as_ref() != Some(source)) + || request + .filter + .created_from + .is_some_and(|from| record.created_at < from) + || request + .filter + .created_to + .is_some_and(|to| record.created_at > to) + { + return false; + } + let content = record.content.to_lowercase(); + tokens.iter().any(|token| content.contains(token)) + }) + .cloned() + .collect(); + matches.sort_by_key(|record| std::cmp::Reverse(record.created_at)); + matches.truncate(request.limit); + Ok(matches) + } + + fn list_records(&self, request: &MemoryListRequest) -> Result { + let limit = request.limit.min(crate::memory::MAX_MEMORY_LIST_PAGE_SIZE); + if limit == 0 { + return Ok(MemoryListPage { + records: Vec::new(), + next_cursor: None, + }); + } + let state = self.lock(); + let mut records = state + .memory + .records + .values() + .filter(|record| { + record.agent_id == request.agent_id + && request.filter.kind.is_none_or(|kind| record.kind == kind) + && request + .filter + .pinned + .is_none_or(|pinned| record.pinned == pinned) + && request + .filter + .source + .as_ref() + .is_none_or(|source| record.source.as_ref() == Some(source)) + && request + .filter + .created_from + .is_none_or(|from| record.created_at >= from) + && request + .filter + .created_to + .is_none_or(|to| record.created_at <= to) + && request + .cursor + .as_ref() + .is_none_or(|cursor| match request.sort { + MemoryListSort::Newest => { + (record.created_at, record.record_id.as_str()) + < (cursor.created_at, cursor.record_id.as_str()) + } + MemoryListSort::Oldest => { + (record.created_at, record.record_id.as_str()) + > (cursor.created_at, cursor.record_id.as_str()) + } + }) + }) + .cloned() + .collect::>(); + records.sort_by(|left, right| match request.sort { + MemoryListSort::Newest => { + (right.created_at, &right.record_id).cmp(&(left.created_at, &left.record_id)) + } + MemoryListSort::Oldest => { + (left.created_at, &left.record_id).cmp(&(right.created_at, &right.record_id)) + } + }); + let has_more = records.len() > limit; + records.truncate(limit); + let next_cursor = if has_more { + records + .last() + .map(|record| crate::memory::MemoryListCursor { + created_at: record.created_at, + record_id: record.record_id.clone(), + }) + } else { + None + }; + Ok(MemoryListPage { + records, + next_cursor, + }) + } + + fn get_record( + &self, + agent_id: &str, + record_id: &str, + ) -> Result, RuntimeError> { + Ok(self + .lock() + .memory + .records + .get(record_id) + .filter(|record| record.agent_id == agent_id) + .cloned()) + } + + fn count_records(&self, agent_id: &str) -> Result { + Ok(self + .lock() + .memory + .records + .values() + .filter(|record| record.agent_id == agent_id) + .count()) + } + + fn delete_records(&self, record_ids: &[String]) -> Result<(), RuntimeError> { + let mut state = self.lock(); + for id in record_ids { + state.memory.records.remove(id); + } + Ok(()) + } + + fn tombstone_records( + &self, + agent_id: &str, + record_ids: &[String], + ) -> Result { + let mut state = self.lock(); + let mut affected = 0usize; + for id in record_ids { + let matches_agent = state + .memory + .records + .get(id) + .is_some_and(|record| record.agent_id == agent_id); + if matches_agent { + state.memory.records.remove(id); + affected += 1; + } + } + Ok(affected) + } + + fn load_agent_memory_cursor( + &self, + agent_id: &str, + ) -> Result, RuntimeError> { + Ok(self.lock().memory.cursors.get(agent_id).cloned()) + } + + fn save_agent_memory_cursor( + &self, + agent_id: &str, + cursor: &MemoryCursor, + ) -> Result<(), RuntimeError> { + self.lock() + .memory + .cursors + .insert(agent_id.to_string(), cursor.clone()); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use crate::memory::{MemoryCursor, MemoryRecord, MemoryRecordKind, MemoryStore}; + + use super::super::VolatileRuntimeStore; + + fn record(id: &str, agent_id: &str, content: &str, created_at: i64) -> MemoryRecord { + MemoryRecord { + record_id: id.to_string(), + agent_id: agent_id.to_string(), + kind: MemoryRecordKind::Episode, + content: content.to_string(), + source_revision: 1, + created_at, + metadata_json: "{}".to_string(), + source: None, + pinned: false, + score: None, + } + } + + #[test] + fn search_is_scoped_to_the_requesting_agent() { + let store = VolatileRuntimeStore::new(); + store + .upsert_records(&[ + record("episode:a:1", "agent-a", "shared phrase alpha", 1), + record("episode:b:1", "agent-b", "shared phrase alpha", 2), + ]) + .expect("seed records"); + + let hits = store + .search_records("agent-a", "alpha", 10) + .expect("search agent-a"); + assert_eq!(hits.len(), 1); + assert_eq!(hits[0].record_id, "episode:a:1"); + } + + #[test] + fn search_ignores_non_searchable_queries() { + let store = VolatileRuntimeStore::new(); + store + .upsert_records(&[record("episode:a:1", "agent-a", "alpha", 1)]) + .expect("seed records"); + + assert!( + store + .search_records("agent-a", "... ---", 10) + .expect("search punctuation-only query") + .is_empty() + ); + } + + #[test] + fn tombstone_only_removes_records_owned_by_the_agent() { + let store = VolatileRuntimeStore::new(); + store + .upsert_records(&[record("episode:a:1", "agent-a", "alpha", 1)]) + .expect("seed records"); + + let affected = store + .tombstone_records("agent-b", &["episode:a:1".to_string()]) + .expect("tombstone with wrong owner"); + assert_eq!(affected, 0); + + let affected = store + .tombstone_records("agent-a", &["episode:a:1".to_string()]) + .expect("tombstone with correct owner"); + assert_eq!(affected, 1); + } + + #[test] + fn memory_cursor_round_trips() { + let store = VolatileRuntimeStore::new(); + assert_eq!( + store + .load_agent_memory_cursor("agent-a") + .expect("load absent cursor"), + None + ); + + let cursor = MemoryCursor { + last_ingested_revision: 7, + }; + store + .save_agent_memory_cursor("agent-a", &cursor) + .expect("save cursor"); + assert_eq!( + store + .load_agent_memory_cursor("agent-a") + .expect("load saved cursor"), + Some(cursor) + ); + } +} diff --git a/vendor/mentra/src/runtime/volatile_store/permission.rs b/vendor/mentra/src/runtime/volatile_store/permission.rs new file mode 100644 index 0000000..9e463e2 --- /dev/null +++ b/vendor/mentra/src/runtime/volatile_store/permission.rs @@ -0,0 +1,189 @@ +use crate::{ + runtime::{PermissionRuleStore, RuntimeError}, + session::{PermissionRuleScope, permission::RememberedRule}, +}; + +use super::VolatileRuntimeStore; + +struct StoredRule { + session_id: String, + project_id: Option, + rule: RememberedRule, +} + +/// Permission rules mirroring the default store's session/project/global +/// scoping in `permission_rules`. +#[derive(Default)] +pub(super) struct PermissionState { + rules: Vec, +} + +impl PermissionRuleStore for VolatileRuntimeStore { + fn save_rules( + &self, + session_id: &str, + project_id: Option<&str>, + rules: &[RememberedRule], + ) -> Result<(), RuntimeError> { + let mut state = self.lock(); + // Only session-scoped rules for this session are replaced; project- + // and global-scoped rules are managed separately and untouched here, + // matching the default store. + state.permissions.rules.retain(|stored| { + !(stored.session_id == session_id && stored.rule.scope == PermissionRuleScope::Session) + }); + for rule in rules { + state.permissions.rules.push(StoredRule { + session_id: session_id.to_string(), + project_id: project_id.map(str::to_string), + rule: rule.clone(), + }); + } + Ok(()) + } + + fn load_rules( + &self, + session_id: &str, + project_id: Option<&str>, + ) -> Result, RuntimeError> { + let state = self.lock(); + Ok(state + .permissions + .rules + .iter() + .filter(|stored| match stored.rule.scope { + PermissionRuleScope::Session => stored.session_id == session_id, + PermissionRuleScope::Project => { + project_id.is_some() && stored.project_id.as_deref() == project_id + } + PermissionRuleScope::Global => true, + }) + .map(|stored| stored.rule.clone()) + .collect()) + } + + fn clear_rules(&self, session_id: &str) -> Result<(), RuntimeError> { + self.lock() + .permissions + .rules + .retain(|stored| stored.session_id != session_id); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use crate::{ + runtime::PermissionRuleStore, + session::{ + PermissionRuleScope, + permission::{RememberedRule, RuleKey}, + }, + }; + + use super::super::VolatileRuntimeStore; + + fn rule(tool_name: &str, allow: bool, scope: PermissionRuleScope) -> RememberedRule { + RememberedRule { + key: RuleKey { + tool_name: tool_name.to_string(), + pattern: None, + }, + allow, + scope, + reason: None, + } + } + + #[test] + fn save_load_clear_round_trip_scoped_by_session_and_project() { + let store = VolatileRuntimeStore::new(); + + store + .save_rules( + "session-a", + None, + &[rule("shell", true, PermissionRuleScope::Session)], + ) + .expect("save session-a rules"); + store + .save_rules( + "session-b", + Some("proj-b"), + &[rule("read", false, PermissionRuleScope::Project)], + ) + .expect("save session-b rules"); + + let loaded_a = store.load_rules("session-a", None).expect("load session-a"); + assert_eq!(loaded_a.len(), 1); + assert_eq!(loaded_a[0].key.tool_name, "shell"); + + let loaded_b = store + .load_rules("session-b", Some("proj-b")) + .expect("load session-b"); + assert_eq!(loaded_b.len(), 1); + assert_eq!(loaded_b[0].key.tool_name, "read"); + + // session-b's project rule does not leak into session-a's load + // without a matching project id. + assert!( + store + .load_rules("session-b", None) + .expect("load session-b without project id") + .is_empty() + ); + + store.clear_rules("session-a").expect("clear session-a"); + assert!( + store + .load_rules("session-a", None) + .expect("load after clear") + .is_empty() + ); + } + + #[test] + fn save_rules_replaces_only_session_scoped_rules() { + let store = VolatileRuntimeStore::new(); + store + .save_rules( + "session-1", + None, + &[rule("shell", true, PermissionRuleScope::Session)], + ) + .expect("save initial"); + store + .save_rules( + "session-1", + None, + &[rule("write", false, PermissionRuleScope::Session)], + ) + .expect("save replacement"); + + let loaded = store.load_rules("session-1", None).expect("load rules"); + assert_eq!(loaded.len(), 1); + assert_eq!(loaded[0].key.tool_name, "write"); + } + + #[test] + fn a_remembered_refusal_keeps_its_reason() { + let store = VolatileRuntimeStore::new(); + let refusal = RememberedRule { + reason: Some("this run does not allow writes".to_string()), + ..rule("write", false, PermissionRuleScope::Session) + }; + + store + .save_rules("session-1", None, &[refusal]) + .expect("save refusal"); + + let loaded = store.load_rules("session-1", None).expect("load refusal"); + assert_eq!(loaded.len(), 1); + assert_eq!( + loaded[0].reason.as_deref(), + Some("this run does not allow writes"), + "the volatile store answers a remembered refusal the same as the persistent one" + ); + } +} diff --git a/vendor/mentra/src/runtime/volatile_store/task.rs b/vendor/mentra/src/runtime/volatile_store/task.rs new file mode 100644 index 0000000..d4fef77 --- /dev/null +++ b/vendor/mentra/src/runtime/volatile_store/task.rs @@ -0,0 +1,154 @@ +use std::{collections::HashMap, path::Path}; + +use crate::runtime::{RuntimeError, TaskItem, TaskStateSnapshot, TaskStore}; + +use super::{VolatileRuntimeStore, path_key}; + +/// Tasks namespaced by the caller-supplied `tasks_dir` path, mirroring the +/// default store's `tasks` table (keyed by the same string). +#[derive(Default)] +pub(super) struct TaskState { + by_namespace: HashMap>, +} + +impl TaskStore for VolatileRuntimeStore { + fn load_tasks(&self, namespace: &Path) -> Result, RuntimeError> { + Ok(self + .lock() + .tasks + .by_namespace + .get(&path_key(namespace)) + .cloned() + .unwrap_or_default()) + } + + fn capture_tasks(&self, namespace: &Path) -> Result { + Ok(TaskStateSnapshot { + tasks: self.load_tasks(namespace)?, + }) + } + + fn restore_tasks( + &self, + namespace: &Path, + snapshot: &TaskStateSnapshot, + ) -> Result<(), RuntimeError> { + self.replace_tasks(namespace, &snapshot.tasks) + } + + fn replace_tasks(&self, namespace: &Path, tasks: &[TaskItem]) -> Result<(), RuntimeError> { + self.lock() + .tasks + .by_namespace + .insert(path_key(namespace), tasks.to_vec()); + Ok(()) + } + + fn mutate( + &self, + namespace: &Path, + mutation: &mut dyn FnMut(&mut Vec) -> Result<(), RuntimeError>, + ) -> Result<(), RuntimeError> { + let mut state = self.lock(); + let key = path_key(namespace); + let mut tasks = state + .tasks + .by_namespace + .get(&key) + .cloned() + .unwrap_or_default(); + mutation(&mut tasks)?; + state.tasks.by_namespace.insert(key, tasks); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use std::path::PathBuf; + + use crate::runtime::{TaskItem, TaskStatus, TaskStore}; + + use super::super::VolatileRuntimeStore; + + fn task(id: u64, subject: &str) -> TaskItem { + TaskItem { + id, + subject: subject.to_string(), + description: String::new(), + status: TaskStatus::Pending, + blocked_by: Vec::new(), + blocks: Vec::new(), + owner: String::new(), + working_directory: None, + } + } + + #[test] + fn load_tasks_reads_own_writes_and_stays_namespaced() { + let store = VolatileRuntimeStore::new(); + let namespace = PathBuf::from("/tmp/does-not-exist/tasks"); + let item = task(1, "write the report"); + + store + .replace_tasks(&namespace, std::slice::from_ref(&item)) + .expect("replace tasks"); + + assert_eq!( + store.load_tasks(&namespace).expect("load tasks"), + vec![item] + ); + assert!( + store + .load_tasks(&PathBuf::from("/tmp/does-not-exist/other")) + .expect("load unrelated namespace") + .is_empty() + ); + } + + #[test] + fn capture_and_restore_round_trip() { + let store = VolatileRuntimeStore::new(); + let namespace = PathBuf::from("/tmp/does-not-exist/tasks-2"); + let item = task(1, "first"); + + store + .replace_tasks(&namespace, std::slice::from_ref(&item)) + .expect("seed"); + let snapshot = store.capture_tasks(&namespace).expect("capture"); + + store.replace_tasks(&namespace, &[]).expect("clear"); + assert!(store.load_tasks(&namespace).expect("load empty").is_empty()); + + store.restore_tasks(&namespace, &snapshot).expect("restore"); + assert_eq!( + store.load_tasks(&namespace).expect("load restored"), + vec![item] + ); + } + + #[test] + fn failed_mutation_does_not_install_partial_changes() { + let store = VolatileRuntimeStore::new(); + let namespace = PathBuf::from("/tmp/does-not-exist/tasks-rollback"); + let item = task(1, "original"); + store + .replace_tasks(&namespace, std::slice::from_ref(&item)) + .expect("seed"); + + let mut mutation = |tasks: &mut Vec| { + tasks[0].subject = "partial".to_string(); + Err(crate::runtime::RuntimeError::InvalidTask( + "reject mutation".to_string(), + )) + }; + store + .mutate(&namespace, &mut mutation) + .expect_err("mutation should fail"); + + assert_eq!( + store.load_tasks(&namespace).expect("load tasks"), + vec![item] + ); + } +} diff --git a/vendor/mentra/src/runtime/volatile_store/team.rs b/vendor/mentra/src/runtime/volatile_store/team.rs new file mode 100644 index 0000000..48d06a2 --- /dev/null +++ b/vendor/mentra/src/runtime/volatile_store/team.rs @@ -0,0 +1,262 @@ +use std::{collections::HashMap, path::Path}; + +use crate::{ + runtime::RuntimeError, + team::{TeamMemberSummary, TeamMessage, TeamProtocolRequestSummary, TeamStore}, +}; + +use super::{DeliveryState, VolatileRuntimeStore, path_key}; + +struct TeamInboxEntry { + team_dir: String, + recipient: String, + message: TeamMessage, + delivery_state: DeliveryState, +} + +/// Team roster, inbox, and protocol-request state namespaced by the +/// caller-supplied `team_dir` path, mirroring the default store's +/// `team_members` / `team_inbox` / `team_requests` tables. +#[derive(Default)] +pub(super) struct TeamState { + members: HashMap<(String, String), TeamMemberSummary>, + inbox: Vec, + requests: HashMap, +} + +impl TeamStore for VolatileRuntimeStore { + fn unread_team_count(&self, team_dir: &Path, agent_name: &str) -> Result { + let state = self.lock(); + let key = path_key(team_dir); + Ok(state + .team + .inbox + .iter() + .filter(|entry| { + entry.team_dir == key + && entry.recipient == agent_name + && entry.delivery_state == DeliveryState::Pending + }) + .count()) + } + + fn load_team_members(&self, team_dir: &Path) -> Result, RuntimeError> { + let state = self.lock(); + let key = path_key(team_dir); + let mut members: Vec<_> = state + .team + .members + .iter() + .filter(|((dir, _), _)| dir == &key) + .map(|(_, summary)| summary.clone()) + .collect(); + members.sort_by(|a, b| a.name.cmp(&b.name)); + Ok(members) + } + + fn upsert_team_member( + &self, + team_dir: &Path, + summary: &TeamMemberSummary, + ) -> Result<(), RuntimeError> { + self.lock() + .team + .members + .insert((path_key(team_dir), summary.name.clone()), summary.clone()); + Ok(()) + } + + fn read_team_inbox( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result, RuntimeError> { + let mut state = self.lock(); + let key = path_key(team_dir); + let mut out = Vec::new(); + for entry in state.team.inbox.iter_mut() { + if entry.team_dir == key + && entry.recipient == agent_name + && entry.delivery_state == DeliveryState::Pending + { + entry.delivery_state = DeliveryState::Inflight; + out.push(entry.message.clone()); + } + } + Ok(out) + } + + fn ack_team_inbox(&self, team_dir: &Path, agent_name: &str) -> Result<(), RuntimeError> { + let mut state = self.lock(); + let key = path_key(team_dir); + for entry in state.team.inbox.iter_mut() { + if entry.team_dir == key + && entry.recipient == agent_name + && entry.delivery_state == DeliveryState::Inflight + { + entry.delivery_state = DeliveryState::Acked; + } + } + Ok(()) + } + + fn requeue_team_inbox(&self, team_dir: &Path, agent_name: &str) -> Result<(), RuntimeError> { + let mut state = self.lock(); + let key = path_key(team_dir); + for entry in state.team.inbox.iter_mut() { + if entry.team_dir == key + && entry.recipient == agent_name + && entry.delivery_state == DeliveryState::Inflight + { + entry.delivery_state = DeliveryState::Pending; + } + } + Ok(()) + } + + fn append_team_message( + &self, + team_dir: &Path, + recipient: &str, + message: &TeamMessage, + ) -> Result<(), RuntimeError> { + self.lock().team.inbox.push(TeamInboxEntry { + team_dir: path_key(team_dir), + recipient: recipient.to_string(), + message: message.clone(), + delivery_state: DeliveryState::Pending, + }); + Ok(()) + } + + fn load_team_requests( + &self, + team_dir: &Path, + ) -> Result, RuntimeError> { + let state = self.lock(); + let key = path_key(team_dir); + let mut requests: Vec<_> = state + .team + .requests + .values() + .filter(|(dir, _)| dir == &key) + .map(|(_, request)| request.clone()) + .collect(); + requests.sort_by(|a, b| { + a.created_at + .cmp(&b.created_at) + .then_with(|| a.request_id.cmp(&b.request_id)) + }); + Ok(requests) + } + + fn upsert_team_request( + &self, + team_dir: &Path, + request: &TeamProtocolRequestSummary, + ) -> Result<(), RuntimeError> { + self.lock().team.requests.insert( + request.request_id.clone(), + (path_key(team_dir), request.clone()), + ); + Ok(()) + } + + fn list_team_agent_names(&self, team_dir: &Path) -> Result, RuntimeError> { + let state = self.lock(); + let key = path_key(team_dir); + let mut names: Vec<_> = state + .agents + .values() + .filter(|record| path_key(&record.config.team.team_dir) == key) + .map(|record| record.name.clone()) + .collect(); + names.sort(); + Ok(names) + } +} + +#[cfg(test)] +mod tests { + use std::path::PathBuf; + + use crate::team::{TeamMessage, TeamMessageKind, TeamStore}; + + use super::super::VolatileRuntimeStore; + + fn message(from: &str, content: &str) -> TeamMessage { + TeamMessage { + kind: TeamMessageKind::Message, + sender: from.to_string(), + content: content.to_string(), + timestamp: 0, + request_id: None, + protocol: None, + approve: None, + } + } + + #[test] + fn team_inbox_round_trips_through_read_ack() { + let store = VolatileRuntimeStore::new(); + let team_dir = PathBuf::from("/tmp/does-not-exist/team"); + + store + .append_team_message(&team_dir, "primary", &message("lead", "hello")) + .expect("append message"); + assert_eq!( + store + .unread_team_count(&team_dir, "primary") + .expect("unread count"), + 1 + ); + + let read = store + .read_team_inbox(&team_dir, "primary") + .expect("read inbox"); + assert_eq!(read.len(), 1); + assert_eq!(read[0].content, "hello"); + + // Read moves messages to in-flight; a second read sees nothing new + // until ack/requeue resolves the in-flight batch. + assert!( + store + .read_team_inbox(&team_dir, "primary") + .expect("second read") + .is_empty() + ); + + store + .ack_team_inbox(&team_dir, "primary") + .expect("ack inbox"); + assert_eq!( + store + .unread_team_count(&team_dir, "primary") + .expect("unread count after ack"), + 0 + ); + } + + #[test] + fn requeue_returns_inflight_messages_to_pending() { + let store = VolatileRuntimeStore::new(); + let team_dir = PathBuf::from("/tmp/does-not-exist/team-2"); + + store + .append_team_message(&team_dir, "primary", &message("lead", "retry me")) + .expect("append message"); + store + .read_team_inbox(&team_dir, "primary") + .expect("read inbox"); + store + .requeue_team_inbox(&team_dir, "primary") + .expect("requeue inbox"); + + assert_eq!( + store + .unread_team_count(&team_dir, "primary") + .expect("unread count after requeue"), + 1 + ); + } +} diff --git a/vendor/mentra/src/session.rs b/vendor/mentra/src/session.rs new file mode 100644 index 0000000..0f0c6e0 --- /dev/null +++ b/vendor/mentra/src/session.rs @@ -0,0 +1,16 @@ +mod event; +mod handle; +pub(crate) mod hooks; +pub(crate) mod mapping; +pub mod permission; +#[cfg(test)] +mod tests; +mod types; + +pub use event::{ + EventSeq, NoticeSeverity, PermissionOutcome, PermissionRuleScope, SessionEvent, TaskKind, + TaskLifecycleStatus, ToolMutability, +}; +pub use handle::{Session, SessionEventReceiver, SessionPermissionHandle, SubagentHandle}; +pub use permission::{PermissionDecision, PermissionRequest, RememberedRule, RuleKey, RuleStore}; +pub use types::{SessionId, SessionMetadata, SessionStatus}; diff --git a/vendor/mentra/src/session/event.rs b/vendor/mentra/src/session/event.rs new file mode 100644 index 0000000..ddea86a --- /dev/null +++ b/vendor/mentra/src/session/event.rs @@ -0,0 +1,170 @@ +use serde::{Deserialize, Serialize}; + +use super::types::SessionId; + +pub type EventSeq = u64; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ToolMutability { + ReadOnly, + Mutating, + Unknown, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum TaskLifecycleStatus { + Spawned, + Running, + Finished, + Failed, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum TaskKind { + Subagent, + BackgroundTask, + Teammate, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum PermissionOutcome { + Allowed, + Denied, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum PermissionRuleScope { + Session, + Project, + Global, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum NoticeSeverity { + Info, + Warning, +} + +/// Events emitted during a session lifecycle. +/// +/// `serde_json::Value` does not implement `Eq`, so the `preview` field in +/// `PermissionRequested` is stored as a JSON `String` to preserve `Eq` +/// derivation on the entire enum. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum SessionEvent { + SessionStarted { + session_id: SessionId, + }, + UserMessage { + text: String, + }, + AssistantTokenDelta { + delta: String, + full_text: String, + }, + AssistantReasoningDelta { + delta: String, + full_text: String, + }, + AssistantMessageCompleted { + text: String, + }, + ToolQueued { + tool_call_id: String, + tool_name: String, + summary: String, + mutability: ToolMutability, + input_json: String, + }, + ToolStarted { + tool_call_id: String, + tool_name: String, + }, + ToolProgress { + tool_call_id: String, + tool_name: String, + progress: String, + }, + ToolCompleted { + tool_call_id: String, + tool_name: String, + summary: String, + is_error: bool, + }, + PermissionRequested { + request_id: String, + tool_call_id: String, + tool_name: String, + description: String, + /// JSON-encoded preview data. Stored as `String` because + /// `serde_json::Value` does not implement `Eq`. + preview: String, + }, + PermissionResolved { + request_id: String, + tool_call_id: String, + tool_name: String, + outcome: PermissionOutcome, + rule_scope: Option, + }, + TaskUpdated { + task_id: String, + kind: TaskKind, + status: TaskLifecycleStatus, + title: String, + detail: Option, + }, + CompactionStarted { + agent_id: String, + }, + CompactionCompleted { + agent_id: String, + replaced_items: usize, + preserved_items: usize, + resulting_transcript_len: usize, + extracted_facts_count: usize, + summary_preview: String, + }, + MemoryUpdated { + agent_id: String, + stored_records: usize, + }, + /// Token usage report after a model response completes. + UsageReport { + agent_id: String, + input_tokens: u64, + output_tokens: u64, + cache_read_tokens: u64, + cache_creation_tokens: u64, + }, + Notice { + severity: NoticeSeverity, + message: String, + }, + RetryAttempt { + agent_id: String, + error_message: String, + attempt: u32, + max_attempts: u32, + next_delay_ms: u64, + }, + Error { + message: String, + recoverable: bool, + }, + /// The session returned to an earlier entry; subsequent turns continue + /// from there along a new path. + Branched { + entry_id: String, + /// How many entries left the active path. They remain in the + /// transcript and stay reachable. + abandoned_entries: usize, + }, +} diff --git a/vendor/mentra/src/session/handle.rs b/vendor/mentra/src/session/handle.rs new file mode 100644 index 0000000..5a6a4cc --- /dev/null +++ b/vendor/mentra/src/session/handle.rs @@ -0,0 +1,690 @@ +use std::sync::Arc; +use std::sync::Mutex as StdMutex; + +use serde::de::DeserializeOwned; +use tokio::sync::broadcast; + +use crate::{ + AgentTranscript, ContentBlock, Message, Role, + agent::{Agent, AgentEvent, AgentEventTapGuard, FinalOutput, TerminalOutputSpec}, + error::RuntimeError, + runtime::{PermissionRuleStore, RunOptions, is_transient_runtime_error}, + session::{ + event::{EventSeq, PermissionOutcome, SessionEvent, TaskKind, TaskLifecycleStatus}, + mapping::{ToolNameIndex, map_agent_event}, + permission::{ + PendingPermissionStore, PermissionDecision, RememberedRule, RuleKey, RuleStore, + }, + types::{SessionId, SessionMetadata, SessionStatus}, + }, + transcript::{EntryId, TranscriptItem}, +}; + +/// Type alias for the receiver end of the session event broadcast channel. +pub type SessionEventReceiver = broadcast::Receiver; + +/// Handle returned from `Session::spawn_subagent` for tracking spawned work. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SubagentHandle { + /// Unique identifier for the spawned task. + pub task_id: String, + /// The subagent's internal agent identifier. + pub agent_id: String, +} + +#[derive(Clone)] +pub struct SessionPermissionHandle { + session_id: SessionId, + project_id: Option, + event_tx: broadcast::Sender, + rule_store: RuleStore, + permission_store: Arc>>>, + pending_permissions: PendingPermissionStore, +} + +impl SessionPermissionHandle { + fn new( + session_id: SessionId, + project_id: Option, + event_tx: broadcast::Sender, + rule_store: RuleStore, + permission_store: Arc>>>, + pending_permissions: PendingPermissionStore, + ) -> Self { + Self { + session_id, + project_id, + event_tx, + rule_store, + permission_store, + pending_permissions, + } + } + + fn set_permission_store(&self, store: Arc) { + let mut slot = self + .permission_store + .lock() + .unwrap_or_else(|e| e.into_inner()); + *slot = Some(store); + } + + fn load_persisted_rules(&self, session_id: &SessionId) -> Result { + let store = self + .permission_store + .lock() + .unwrap_or_else(|e| e.into_inner()) + .clone(); + let Some(store) = store else { + return Ok(0); + }; + let rules = store.load_rules(session_id.as_str(), self.project_id.as_deref())?; + let count = rules.len(); + for rule in rules { + self.rule_store.add_rule(rule); + } + Ok(count) + } + + pub fn resolve_permission( + &self, + request_id: &str, + decision: PermissionDecision, + ) -> Result<(), RuntimeError> { + let entry = self.pending_permissions.remove(request_id).ok_or_else(|| { + RuntimeError::OperationDenied(format!( + "no pending permission with request_id '{request_id}'" + )) + })?; + + let outcome = if decision.allow { + PermissionOutcome::Allowed + } else { + PermissionOutcome::Denied + }; + + if let Some(scope) = decision.remember_as { + self.rule_store.add_rule(RememberedRule { + key: RuleKey { + tool_name: entry.tool_name.clone(), + pattern: None, + }, + allow: decision.allow, + scope, + // Only a refusal has anything left to say: this call's reason + // reaches the model as its tool result, and every later call + // the rule answers has nothing else to read. + reason: if decision.allow { + None + } else { + decision.reason.clone() + }, + }); + + let store = self + .permission_store + .lock() + .unwrap_or_else(|e| e.into_inner()) + .clone(); + if let Some(store) = store { + let all_rules = self.rule_store.rules(); + store.save_rules( + self.session_id.as_str(), + self.project_id.as_deref(), + &all_rules, + )?; + } + } + + let _ = self.event_tx.send(SessionEvent::PermissionResolved { + request_id: request_id.to_owned(), + tool_call_id: entry.tool_call_id, + tool_name: entry.tool_name, + outcome, + rule_scope: decision.remember_as, + }); + + let _ = entry.sender.send(decision); + Ok(()) + } + + pub(crate) fn remembered_rules(&self) -> Vec { + self.rule_store.rules() + } + + pub(crate) fn rule_store(&self) -> &RuleStore { + &self.rule_store + } +} + +/// A `Session` wraps an [`Agent`] with session-level metadata and a broadcast +/// event channel that emits [`SessionEvent`] values for UI consumption. +pub struct Session { + id: SessionId, + metadata: SessionMetadata, + agent: Agent, + event_tx: broadcast::Sender, + next_seq: EventSeq, + /// Shared with the per-turn event tap so a tool call queued in one turn + /// still resolves its name when the result arrives in another. + tool_names: Arc>, + #[allow(dead_code)] + pub(crate) pending_permissions: PendingPermissionStore, + permission_handle: SessionPermissionHandle, +} + +impl Session { + /// Creates a new session wrapping the given agent. + #[allow(dead_code)] + pub(crate) fn new(id: SessionId, metadata: SessionMetadata, agent: Agent) -> Self { + let (event_tx, _) = broadcast::channel(512); + Self::new_with_parts( + id, + metadata, + agent, + event_tx, + RuleStore::new(), + PendingPermissionStore::new(), + None, + ) + } + + pub(crate) fn new_with_parts( + id: SessionId, + metadata: SessionMetadata, + agent: Agent, + event_tx: broadcast::Sender, + rule_store: RuleStore, + pending_permissions: PendingPermissionStore, + project_id: Option, + ) -> Self { + let permission_store = Arc::new(StdMutex::new(None)); + let permission_handle = SessionPermissionHandle::new( + id.clone(), + project_id, + event_tx.clone(), + rule_store.clone(), + permission_store.clone(), + pending_permissions.clone(), + ); + Self { + id, + metadata, + agent, + event_tx, + next_seq: 0, + tool_names: Arc::new(StdMutex::new(ToolNameIndex::default())), + pending_permissions, + permission_handle, + } + } + + /// Attaches a persistent permission rule store to this session. + /// + /// When set, remembered rules are saved to the store on each decision and + /// can be loaded on session resume via [`load_persisted_rules`](Self::load_persisted_rules). + pub fn set_permission_store(&mut self, store: Arc) { + self.permission_handle.set_permission_store(store); + } + + /// Loads persisted permission rules from the attached store into the + /// in-memory [`RuleStore`]. + /// + /// This is typically called during session resume to restore rules that were + /// persisted in a prior session run. Returns the number of rules loaded. + pub fn load_persisted_rules(&mut self) -> Result { + self.permission_handle.load_persisted_rules(&self.id) + } + + /// Returns the session identifier. + pub fn id(&self) -> &SessionId { + &self.id + } + + /// Returns the session metadata (title, model, status, turn count, timestamps). + pub fn metadata(&self) -> &SessionMetadata { + &self.metadata + } + + /// Updates the live session model and persists the new setting so future + /// resumes observe the same model. + pub fn set_model(&mut self, model: crate::ModelInfo) -> Result<(), RuntimeError> { + self.agent.set_model(model.clone())?; + self.metadata.model = model.id; + self.metadata.updated_at = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + Ok(()) + } + + /// Updates the reasoning options for the live session's agent and persists the + /// new setting (mirrors [`set_model`](Self::set_model)). Lets a caller run + /// per-phase tiering — e.g. a low effort while gathering, then a higher effort + /// for a final synthesis turn — on the same session. + pub fn set_reasoning( + &mut self, + reasoning: Option, + ) -> Result<(), RuntimeError> { + self.agent.set_reasoning(reasoning) + } + + /// Returns the underlying agent identifier. + pub fn agent_id(&self) -> &str { + self.agent.id() + } + + /// Returns the session display name (same as the agent name). + pub fn name(&self) -> &str { + self.agent.name() + } + + /// Subscribes to the session event stream. + pub fn subscribe(&self) -> SessionEventReceiver { + self.event_tx.subscribe() + } + + pub fn permission_handle(&self) -> SessionPermissionHandle { + self.permission_handle.clone() + } + + /// Submits a user turn, runs the agent, emits session events, and returns + /// the assistant response message. + pub async fn append_turn( + &mut self, + content: Vec, + ) -> Result { + self.append_turn_with_options(content, RunOptions::default()) + .await + } + + /// Submits a user turn with explicit execution limits and cancellation + /// settings. + /// + /// The session-level counterpart to [`Agent::run`]. A host that drives a + /// conversation through a `Session` — for the event stream and the + /// permission handle — needs the same control over a turn that + /// [`Agent::run`] gives, without dropping to the agent and losing both. + /// Cancelling through [`RunOptions::cancellation`] fails the turn and rolls + /// it back; [`RunOptions::stop`] ends it gracefully at the next round + /// boundary, keeping the committed transcript. + pub async fn append_turn_with_options( + &mut self, + content: Vec, + options: RunOptions, + ) -> Result { + let user_text = extract_user_text(&content); + self.emit(SessionEvent::UserMessage { text: user_text }); + + let turn = self.begin_turn(); + let result = self.agent.run(content, options).await; + self.finish_turn(turn, result) + } + + /// Submits a user turn that must end in a typed value, and returns that + /// value together with the tool-result message that carried it. + /// + /// The session-level counterpart to [`Agent::run_to_output`], as + /// [`append_turn_with_options`](Self::append_turn_with_options) is to + /// [`Agent::run`]. A host that drives a conversation through a `Session` — + /// for the event stream and the permission handle — needs a typed final + /// answer without dropping to the agent and losing both. + /// + /// Whether the turn may work on its way to that answer or only shape what + /// the conversation already holds is the spec's to say, through + /// [`TerminalOutputSpec::with_tools`]. A working typed turn is where a + /// `Session` earns the most: its tool calls go to the same event stream + /// and its permission requests to the same handle as any other turn's, so + /// a host gets one turn that reads, asks, and answers in a declared shape + /// without giving up either. + /// + /// The turn announces itself on the stream exactly as every other turn + /// does: a [`SessionEvent::UserMessage`] going in, whatever the agent + /// emits while it runs, and on success one + /// [`SessionEvent::AssistantMessageCompleted`] carrying the text of the + /// turn's final assistant message. For a typed turn that is whatever prose + /// the model wrote alongside the terminal tool call, which is often + /// nothing. The value itself is deliberately not put there: it already + /// reaches the stream as the terminal tool's + /// [`ToolQueued`](SessionEvent::ToolQueued) input and + /// [`ToolCompleted`](SessionEvent::ToolCompleted) summary, and a client + /// that reads `AssistantMessageCompleted` as "what the assistant said" + /// would render a tool payload as prose. Failure reports the same way any + /// turn does — [`SessionEvent::Error`] and [`SessionStatus::Failed`]. + /// + /// One asymmetry with a plain turn is worth knowing: a value that does not + /// deserialize into `T` fails *after* the agent committed the exchange, so + /// the transcript holds the terminal call and its result even though this + /// returns `Err`. The turn counter still does not move, as for any failed + /// turn. + pub async fn append_turn_to_output( + &mut self, + content: Vec, + options: RunOptions, + spec: TerminalOutputSpec, + ) -> Result, RuntimeError> { + let user_text = extract_user_text(&content); + self.emit(SessionEvent::UserMessage { text: user_text }); + + let turn = self.begin_turn(); + let result = self.agent.run_to_output::(content, options, spec).await; + self.finish_turn(turn, result) + } + + /// Returns the agent's canonical transcript for UI reconstruction. + pub fn replay(&self) -> &AgentTranscript { + self.agent.transcript() + } + + /// The entry the next turn will continue from. + pub fn leaf(&self) -> Option<&EntryId> { + self.agent.leaf() + } + + /// Returns to an earlier entry so the next turn takes a different path. + /// + /// This is how "undo that exchange and try something else" works without + /// starting a new session: the entries after `entry` leave the active + /// path but stay in the transcript, so the abandoned line of work is + /// still addressable through [`children`](Self::children). Emits + /// [`SessionEvent::Branched`] with the number of entries that moved. + pub fn branch_from(&mut self, entry: &EntryId) -> Result { + let moved = self.agent.branch_from(entry)?; + self.emit(SessionEvent::Branched { + entry_id: entry.to_string(), + abandoned_entries: moved, + }); + self.touch_updated_at(); + Ok(moved) + } + + /// The entries recorded as continuing from `entry`, in creation order. + /// More than one means the conversation branched there. + pub fn children(&self, entry: &EntryId) -> Vec<&TranscriptItem> { + self.agent.children(entry) + } + + /// Resumes the agent from an interrupted or failed state, emitting session + /// events as the turn runs. + pub async fn resume_turn(&mut self) -> Result { + self.resume_turn_with_options(RunOptions::default()).await + } + + /// Resumes an interrupted or failed turn with explicit execution limits and + /// cancellation settings. + pub async fn resume_turn_with_options( + &mut self, + options: RunOptions, + ) -> Result { + let turn = self.begin_turn(); + let result = self.agent.resume_with_options(options).await; + self.finish_turn(turn, result) + } + + /// Opens a turn: marks the session active and starts forwarding agent + /// events onto the session stream. + /// + /// Every turn opens here and closes in [`finish_turn`](Self::finish_turn) — + /// started by a prompt, resumed after a failure, or run to a typed output — + /// so no turn can report itself differently from the others. + fn begin_turn(&mut self) -> TurnGuard { + self.update_status(SessionStatus::Active); + let (event_tap, forwarded_seq) = self.install_agent_event_forwarder(); + TurnGuard { + event_tap, + forwarded_seq, + } + } + + /// Closes a turn opened by [`begin_turn`](Self::begin_turn): stops the + /// forwarder, emits the terminal event, and settles the status, the turn + /// counter, and `updated_at`. + fn finish_turn( + &mut self, + turn: TurnGuard, + result: Result, + ) -> Result { + let TurnGuard { + event_tap, + forwarded_seq, + } = turn; + drop(event_tap); + self.sync_forwarded_seq(&forwarded_seq); + + match result { + Ok(outcome) => { + let text = outcome.completion_text(self.agent.history()); + self.emit(SessionEvent::AssistantMessageCompleted { text }); + self.metadata.turn_count += 1; + self.update_status(SessionStatus::Idle); + self.touch_updated_at(); + Ok(outcome) + } + Err(error) => { + let recoverable = is_transient_runtime_error(&error); + self.emit(SessionEvent::Error { + message: error.to_string(), + recoverable, + }); + self.update_status(SessionStatus::Failed(error.to_string())); + self.touch_updated_at(); + Err(error) + } + } + } + + /// Returns the committed message history. + pub fn history(&self) -> &[Message] { + self.agent.history() + } + + /// Emits the initial `SessionStarted` event. Used by `Runtime::create_session`. + pub(crate) fn emit_started(&mut self, event: SessionEvent) { + self.emit(event); + } + + /// Resolves a pending permission request with the given decision. + /// + /// If `remember_as` is set on the decision, the rule is stored in the + /// session's [`RuleStore`]. A [`SessionEvent::PermissionResolved`] event is + /// emitted and the decision is sent back to the waiting caller via oneshot. + pub fn resolve_permission( + &self, + request_id: &str, + decision: PermissionDecision, + ) -> Result<(), RuntimeError> { + self.permission_handle + .resolve_permission(request_id, decision) + } + + /// Returns all remembered permission rules for this session. + pub fn remembered_rules(&self) -> Vec { + self.permission_handle.remembered_rules() + } + + /// Returns a reference to the session's rule store. + pub fn rule_store(&self) -> &RuleStore { + self.permission_handle.rule_store() + } + + /// Returns summaries of all teammates registered with this session's agent. + pub fn list_teammates(&self) -> Vec { + self.agent.watch_snapshot().borrow().teammates.clone() + } + + /// Returns summaries of all active or recently completed subagents. + pub fn active_subagents(&self) -> Vec { + self.agent.watch_snapshot().borrow().subagents.clone() + } + + /// Spawns a disposable subagent in the background and returns a handle for tracking it. + /// + /// The subagent is registered with the parent agent, a `SubagentSpawned` event is emitted, + /// and the subagent runs its prompt in a detached `tokio::spawn`. When it completes, a + /// `SessionEvent::TaskUpdated` event is broadcast with the final status. + /// + /// The subagent runs on [`RunOptions::default`]: this is a host-initiated + /// spawn with no session turn necessarily in flight, so there is no parent + /// run whose bounds it could inherit. A host that wants the subagent to + /// share a turn's cancellation and token accounting passes that turn's + /// [`RunOptions::child`] to + /// [`spawn_subagent_with_options`](Self::spawn_subagent_with_options) + /// instead. This is the opposite of the model-facing `task` intrinsic, + /// which always inherits because it can only run *inside* a parent run. + pub async fn spawn_subagent( + &mut self, + name: &str, + prompt: &str, + ) -> Result { + self.spawn_subagent_with_options(name, prompt, RunOptions::default()) + .await + } + + /// Spawns a disposable subagent that runs on caller-supplied `options`. + /// + /// Pass a turn's [`RunOptions::child`] to put the subagent under that + /// turn's cancellation, stop, deadline, and shared token accounting. The + /// subagent is detached, so those bounds are the only thing tying its + /// lifetime to the turn's. + pub async fn spawn_subagent_with_options( + &mut self, + name: &str, + prompt: &str, + options: RunOptions, + ) -> Result { + let mut subagent = self.agent.spawn_subagent()?; + let agent_id = subagent.id().to_string(); + let summary = self.agent.register_subagent(&subagent); + + self.agent.emit_event(AgentEvent::SubagentSpawned { + agent: summary.clone(), + }); + + let handle = SubagentHandle { + task_id: agent_id.clone(), + agent_id: agent_id.clone(), + }; + + let event_tx = self.event_tx.clone(); + let task_name = name.to_string(); + let prompt_text = prompt.to_string(); + + tokio::spawn(async move { + let result = subagent + .run(vec![ContentBlock::Text { text: prompt_text }], options) + .await; + + let (status, detail) = match &result { + Ok(msg) => (TaskLifecycleStatus::Finished, Some(msg.text())), + Err(e) => (TaskLifecycleStatus::Failed, Some(e.to_string())), + }; + + let _ = event_tx.send(SessionEvent::TaskUpdated { + task_id: agent_id, + kind: TaskKind::Subagent, + status, + title: task_name, + detail, + }); + }); + + Ok(handle) + } + + // -- internal helpers -- + + fn emit(&mut self, event: SessionEvent) { + // Ignore send errors — there may be no active subscribers. + let _ = self.event_tx.send(event); + self.next_seq += 1; + } + + fn install_agent_event_forwarder(&mut self) -> (AgentEventTapGuard, Arc>) { + let event_tx = self.event_tx.clone(); + let next_seq = Arc::new(StdMutex::new(self.next_seq)); + let next_seq_for_tap = Arc::clone(&next_seq); + let tool_names = Arc::clone(&self.tool_names); + let event_tap = self.agent.register_event_tap(move |agent_event| { + let mut seq = next_seq_for_tap + .lock() + .unwrap_or_else(|error| error.into_inner()); + let mut names = tool_names.lock().unwrap_or_else(|error| error.into_inner()); + let mapped = map_agent_event(agent_event, &mut seq, &mut names); + for (_seq, session_event) in mapped { + let _ = event_tx.send(session_event); + } + }); + (event_tap, next_seq) + } + + fn sync_forwarded_seq(&mut self, next_seq: &Arc>) { + self.next_seq = *next_seq.lock().unwrap_or_else(|error| error.into_inner()); + } + + fn update_status(&mut self, status: SessionStatus) { + self.metadata.status = status; + } + + fn touch_updated_at(&mut self) { + self.metadata.updated_at = unix_now(); + } +} + +/// The bookkeeping a turn holds open while it runs: the agent-event tap that +/// forwards onto the session stream, and the sequence counter the tap advances +/// behind it. +#[must_use = "a turn opened with `begin_turn` must be closed with `finish_turn`"] +struct TurnGuard { + event_tap: AgentEventTapGuard, + forwarded_seq: Arc>, +} + +/// What a successful turn puts in its terminal +/// [`SessionEvent::AssistantMessageCompleted`]. +/// +/// Both shapes a turn can return resolve to one rule — the text of the turn's +/// final assistant message — so the stream reads the same whether the turn +/// returned that message itself or a typed value extracted from the tool +/// result that followed it. +trait TurnOutcome { + fn completion_text(&self, history: &[Message]) -> String; +} + +impl TurnOutcome for Message { + /// A prompted or resumed turn returns the final assistant message itself. + fn completion_text(&self, _history: &[Message]) -> String { + self.text() + } +} + +impl TurnOutcome for FinalOutput { + /// A typed turn ends on a tool-result message, so its final assistant + /// message is the one that carried the terminal call — the last assistant + /// message in the committed history. + fn completion_text(&self, history: &[Message]) -> String { + history + .iter() + .rev() + .find(|message| message.role == Role::Assistant) + .map(Message::text) + .unwrap_or_default() + } +} + +fn extract_user_text(content: &[ContentBlock]) -> String { + content + .iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text.as_str()), + _ => None, + }) + .collect::>() + .join("") +} + +fn unix_now() -> u64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} diff --git a/vendor/mentra/src/session/hooks.rs b/vendor/mentra/src/session/hooks.rs new file mode 100644 index 0000000..18aebdc --- /dev/null +++ b/vendor/mentra/src/session/hooks.rs @@ -0,0 +1,206 @@ +use tokio::sync::broadcast; + +use crate::{ + runtime::{AuditStore, RuntimeError, RuntimeHook, RuntimeHookEvent}, + session::event::{NoticeSeverity, SessionEvent}, +}; + +/// A [`RuntimeHook`] that forwards memory ingest events into the session event +/// broadcast channel as [`SessionEvent::MemoryUpdated`] or +/// [`SessionEvent::Notice`] events. +/// +/// This hook is the bridge between the low-level runtime hook system and the +/// session-level event stream consumed by UI layers. +#[allow(dead_code)] // Session-level memory hook wiring is staged separately from the hook itself. +pub(crate) struct SessionHookBridge { + tx: broadcast::Sender, +} + +#[allow(dead_code)] // Constructor is kept alongside the bridge until runtime/session wiring lands. +impl SessionHookBridge { + /// Creates a new bridge that sends session events into the given sender. + pub(crate) fn new(tx: broadcast::Sender) -> Self { + Self { tx } + } +} + +impl RuntimeHook for SessionHookBridge { + fn on_event( + &self, + _store: &dyn AuditStore, + event: &RuntimeHookEvent, + ) -> Result<(), RuntimeError> { + match event { + RuntimeHookEvent::MemoryIngestFinished { + success: true, + stored_records, + agent_id, + .. + } => { + // Ignore send errors — there may be no active subscribers. + let _ = self.tx.send(SessionEvent::MemoryUpdated { + agent_id: agent_id.clone(), + stored_records: *stored_records, + }); + } + RuntimeHookEvent::MemoryIngestFinished { + success: false, + agent_id, + error, + .. + } => { + let message = error.as_deref().unwrap_or("memory ingest failed"); + let _ = self.tx.send(SessionEvent::Notice { + severity: NoticeSeverity::Warning, + message: format!("agent '{agent_id}': {message}"), + }); + } + _ => {} + } + + Ok(()) + } +} + +#[cfg(test)] +#[expect( + clippy::unwrap_used, + reason = "tests use unwrap to fail immediately on invalid fixtures" +)] +mod tests { + use tokio::sync::broadcast; + + use super::*; + use crate::runtime::RuntimeHookEvent; + use crate::session::event::SessionEvent; + + struct NoopAuditStore; + + impl crate::runtime::AuditStore for NoopAuditStore { + fn record_audit_event( + &self, + _scope: &str, + _event_type: &str, + _payload: serde_json::Value, + ) -> Result<(), RuntimeError> { + Ok(()) + } + } + + fn make_bridge() -> (SessionHookBridge, broadcast::Receiver) { + let (tx, rx) = broadcast::channel(16); + (SessionHookBridge::new(tx), rx) + } + + #[test] + fn memory_ingest_finished_success_emits_memory_updated() { + let (bridge, mut rx) = make_bridge(); + let store = NoopAuditStore; + + let event = RuntimeHookEvent::MemoryIngestFinished { + agent_id: "agent-1".to_string(), + source_revision: 42, + success: true, + stored_records: 7, + error: None, + }; + + bridge.on_event(&store, &event).unwrap(); + + let received = rx.try_recv().unwrap(); + assert!( + matches!( + &received, + SessionEvent::MemoryUpdated { agent_id, stored_records } + if agent_id == "agent-1" && *stored_records == 7 + ), + "Expected MemoryUpdated, got: {received:?}" + ); + } + + #[test] + fn memory_ingest_finished_failure_emits_notice_warning() { + let (bridge, mut rx) = make_bridge(); + let store = NoopAuditStore; + + let event = RuntimeHookEvent::MemoryIngestFinished { + agent_id: "agent-2".to_string(), + source_revision: 1, + success: false, + stored_records: 0, + error: Some("disk full".to_string()), + }; + + bridge.on_event(&store, &event).unwrap(); + + let received = rx.try_recv().unwrap(); + assert!( + matches!( + &received, + SessionEvent::Notice { severity: NoticeSeverity::Warning, message } + if message.contains("agent-2") && message.contains("disk full") + ), + "Expected Warning Notice, got: {received:?}" + ); + } + + #[test] + fn memory_ingest_failure_without_error_uses_fallback_message() { + let (bridge, mut rx) = make_bridge(); + let store = NoopAuditStore; + + let event = RuntimeHookEvent::MemoryIngestFinished { + agent_id: "agent-3".to_string(), + source_revision: 1, + success: false, + stored_records: 0, + error: None, + }; + + bridge.on_event(&store, &event).unwrap(); + + let received = rx.try_recv().unwrap(); + assert!( + matches!( + &received, + SessionEvent::Notice { severity: NoticeSeverity::Warning, message } + if message.contains("memory ingest failed") + ), + "Expected fallback message, got: {received:?}" + ); + } + + #[test] + fn non_memory_events_are_silently_ignored() { + let (bridge, mut rx) = make_bridge(); + let store = NoopAuditStore; + + let event = RuntimeHookEvent::RunAborted { + agent_id: "agent-1".to_string(), + reason: "timeout".to_string(), + }; + + bridge.on_event(&store, &event).unwrap(); + + assert!( + rx.try_recv().is_err(), + "Expected no event emitted for non-memory hook events" + ); + } + + #[test] + fn memory_updated_event_serialization_roundtrip() { + let event = SessionEvent::MemoryUpdated { + agent_id: "agent-1".to_string(), + stored_records: 12, + }; + + let json = serde_json::to_value(&event).unwrap(); + assert_eq!(json["type"], "memory_updated"); + assert_eq!(json["agent_id"], "agent-1"); + assert_eq!(json["stored_records"], 12); + + let deserialized: SessionEvent = serde_json::from_value(json).unwrap(); + assert_eq!(event, deserialized); + } +} diff --git a/vendor/mentra/src/session/mapping.rs b/vendor/mentra/src/session/mapping.rs new file mode 100644 index 0000000..da2fe52 --- /dev/null +++ b/vendor/mentra/src/session/mapping.rs @@ -0,0 +1,945 @@ +use std::collections::HashMap; + +use crate::{ + ContentBlock, + agent::{AgentEvent, CompactionDetails, SpawnedAgentStatus, SpawnedAgentSummary}, + background::{BackgroundTaskStatus, BackgroundTaskSummary}, + session::event::{EventSeq, SessionEvent, TaskKind, TaskLifecycleStatus, ToolMutability}, + team::{TeamMemberStatus, TeamMemberSummary}, + tool::{ToolExecutionCategory, ToolSideEffectLevel}, +}; + +/// Remembers the tool name of each in-flight call. +/// +/// A call announces its name when it is queued and when it starts, but the +/// result arrives as a [`ContentBlock::ToolResult`], which carries only +/// `tool_use_id`. Without this, completion — the event a client most wants to +/// attribute, since it is where failures surface — would be the one point in +/// the lifecycle that cannot name its tool. +/// +/// The index belongs to the session rather than to a turn, so a call queued +/// before an interruption still resolves when the turn resumes. +#[derive(Debug, Default)] +pub(crate) struct ToolNameIndex { + names: HashMap, +} + +impl ToolNameIndex { + fn remember(&mut self, id: &str, name: &str) { + if id.is_empty() || name.is_empty() { + return; + } + self.names.insert(id.to_string(), name.to_string()); + } + + /// Resolves a call and forgets it. A call completes once, so keeping the + /// entry would grow the map for the life of the session. + fn resolve(&mut self, id: &str) -> Option { + self.names.remove(id) + } +} + +/// Maps an `AgentEvent` to zero or more `SessionEvent` values. +/// +/// Some agent events map one-to-one, others produce multiple session events +/// (e.g. compaction), and some are intentionally silenced at the session layer. +pub(crate) fn map_agent_event( + event: &AgentEvent, + seq: &mut EventSeq, + tool_names: &mut ToolNameIndex, +) -> Vec<(EventSeq, SessionEvent)> { + let mut out = Vec::new(); + + let mapped = map_event_inner(event, tool_names); + for session_event in mapped { + let current_seq = *seq; + *seq += 1; + out.push((current_seq, session_event)); + } + + out +} + +fn map_event_inner(event: &AgentEvent, tool_names: &mut ToolNameIndex) -> Vec { + match event { + AgentEvent::TextDelta { delta, full_text } => { + vec![SessionEvent::AssistantTokenDelta { + delta: delta.clone(), + full_text: full_text.clone(), + }] + } + + AgentEvent::ReasoningDelta { delta, full_text } => { + vec![SessionEvent::AssistantReasoningDelta { + delta: delta.clone(), + full_text: full_text.clone(), + }] + } + + AgentEvent::ToolUseReady { call, .. } => { + tool_names.remember(&call.id, &call.name); + let input_str = call.input.to_string(); + let summary = derive_tool_summary(&call.name, &input_str); + vec![SessionEvent::ToolQueued { + tool_call_id: call.id.clone(), + tool_name: call.name.clone(), + summary, + mutability: ToolMutability::Unknown, + input_json: input_str, + }] + } + + AgentEvent::ToolExecutionStarted { call } => { + // Also remembered here: a subscriber is not guaranteed to have + // seen the queue event, and a call can start without one. + tool_names.remember(&call.id, &call.name); + vec![SessionEvent::ToolStarted { + tool_call_id: call.id.clone(), + tool_name: call.name.clone(), + }] + } + + AgentEvent::ToolExecutionFinished { result } => map_tool_result(result, tool_names), + + AgentEvent::ToolExecutionProgress { id, name, progress } => { + vec![SessionEvent::ToolProgress { + tool_call_id: id.clone(), + tool_name: name.clone(), + progress: progress.clone(), + }] + } + + AgentEvent::ContextCompacted { details } => map_compaction(details), + + AgentEvent::SubagentSpawned { agent } => map_subagent(agent, TaskLifecycleStatus::Spawned), + AgentEvent::SubagentFinished { agent } => map_subagent_finished(agent), + + AgentEvent::BackgroundTaskStarted { task } => { + map_background_task(task, TaskLifecycleStatus::Running) + } + AgentEvent::BackgroundTaskFinished { task } => map_background_task_finished(task), + + AgentEvent::TeammateSpawned { teammate } => { + map_teammate(teammate, TaskLifecycleStatus::Spawned) + } + AgentEvent::TeammateUpdated { teammate } => map_teammate_updated(teammate), + + AgentEvent::RetryAttempt { + agent_id, + error_message, + attempt, + max_attempts, + next_delay_ms, + } => vec![SessionEvent::RetryAttempt { + agent_id: agent_id.clone(), + error_message: error_message.clone(), + attempt: *attempt, + max_attempts: *max_attempts, + next_delay_ms: *next_delay_ms, + }], + + AgentEvent::UsageReport { + input_tokens, + output_tokens, + cache_read_tokens, + cache_creation_tokens, + } => vec![SessionEvent::UsageReport { + agent_id: String::new(), + input_tokens: *input_tokens, + output_tokens: *output_tokens, + cache_read_tokens: *cache_read_tokens, + cache_creation_tokens: *cache_creation_tokens, + }], + + // Events handled at Session level or intentionally silent at session layer. + AgentEvent::AssistantMessageCommitted { .. } + | AgentEvent::RunStarted + | AgentEvent::RunFinished + | AgentEvent::RunFailed { .. } + | AgentEvent::ToolUseUpdated { .. } + | AgentEvent::TeamProtocolRequested { .. } + | AgentEvent::TeamProtocolResolved { .. } + | AgentEvent::TeamInboxUpdated { .. } => Vec::new(), + } +} + +fn map_tool_result(block: &ContentBlock, tool_names: &mut ToolNameIndex) -> Vec { + if let ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } = block + { + let summary = truncate_input_summary(&content.to_display_string(), 200); + vec![SessionEvent::ToolCompleted { + tool_call_id: tool_use_id.clone(), + // Empty only for a result whose call this session never saw, which + // the index cannot help with — not for every call, as before. + tool_name: tool_names.resolve(tool_use_id).unwrap_or_default(), + summary, + is_error: *is_error, + }] + } else { + Vec::new() + } +} + +fn map_compaction(details: &CompactionDetails) -> Vec { + vec![ + SessionEvent::CompactionStarted { + agent_id: details.agent_id.clone(), + }, + SessionEvent::CompactionCompleted { + agent_id: details.agent_id.clone(), + replaced_items: details.replaced_items, + preserved_items: details.preserved_items, + resulting_transcript_len: details.resulting_transcript_len, + extracted_facts_count: details.extracted_facts_count, + summary_preview: details.summary_preview.clone(), + }, + ] +} + +fn map_subagent(agent: &SpawnedAgentSummary, status: TaskLifecycleStatus) -> Vec { + vec![SessionEvent::TaskUpdated { + task_id: agent.id.clone(), + kind: TaskKind::Subagent, + status, + title: agent.name.clone(), + detail: None, + }] +} + +fn map_subagent_finished(agent: &SpawnedAgentSummary) -> Vec { + let status = match &agent.status { + SpawnedAgentStatus::Finished => TaskLifecycleStatus::Finished, + SpawnedAgentStatus::Failed(_) => TaskLifecycleStatus::Failed, + SpawnedAgentStatus::Running => TaskLifecycleStatus::Running, + }; + let detail = match &agent.status { + SpawnedAgentStatus::Failed(msg) => Some(msg.clone()), + _ => None, + }; + vec![SessionEvent::TaskUpdated { + task_id: agent.id.clone(), + kind: TaskKind::Subagent, + status, + title: agent.name.clone(), + detail, + }] +} + +fn map_background_task( + task: &BackgroundTaskSummary, + status: TaskLifecycleStatus, +) -> Vec { + vec![SessionEvent::TaskUpdated { + task_id: task.id.clone(), + kind: TaskKind::BackgroundTask, + status, + title: task.command.clone(), + detail: task.output_preview.clone(), + }] +} + +fn map_background_task_finished(task: &BackgroundTaskSummary) -> Vec { + let status = match task.status { + BackgroundTaskStatus::Finished => TaskLifecycleStatus::Finished, + BackgroundTaskStatus::Failed | BackgroundTaskStatus::Interrupted => { + TaskLifecycleStatus::Failed + } + BackgroundTaskStatus::Running => TaskLifecycleStatus::Running, + }; + vec![SessionEvent::TaskUpdated { + task_id: task.id.clone(), + kind: TaskKind::BackgroundTask, + status, + title: task.command.clone(), + detail: task.output_preview.clone(), + }] +} + +fn map_teammate(teammate: &TeamMemberSummary, status: TaskLifecycleStatus) -> Vec { + vec![SessionEvent::TaskUpdated { + task_id: teammate.id.clone(), + kind: TaskKind::Teammate, + status, + title: teammate.name.clone(), + detail: Some(teammate.role.clone()), + }] +} + +fn map_teammate_updated(teammate: &TeamMemberSummary) -> Vec { + let status = match &teammate.status { + TeamMemberStatus::Idle | TeamMemberStatus::Working => TaskLifecycleStatus::Running, + TeamMemberStatus::Shutdown => TaskLifecycleStatus::Finished, + TeamMemberStatus::Failed(_) => TaskLifecycleStatus::Failed, + }; + let detail = match &teammate.status { + TeamMemberStatus::Failed(msg) => Some(msg.clone()), + _ => Some(teammate.role.clone()), + }; + vec![SessionEvent::TaskUpdated { + task_id: teammate.id.clone(), + kind: TaskKind::Teammate, + status, + title: teammate.name.clone(), + detail, + }] +} + +#[allow(dead_code)] // exposed for session-handle enrichment in upcoming tasks +pub(crate) fn classify_mutability( + side_effect_level: ToolSideEffectLevel, + execution_category: ToolExecutionCategory, +) -> ToolMutability { + match (side_effect_level, execution_category) { + (ToolSideEffectLevel::None, _) => ToolMutability::ReadOnly, + (_, ToolExecutionCategory::ReadOnlyParallel) => ToolMutability::ReadOnly, + _ => ToolMutability::Mutating, + } +} + +pub(crate) fn derive_tool_summary(tool_name: &str, input_json: &str) -> String { + if let Ok(value) = serde_json::from_str::(input_json) { + if let Some(command) = value.get("command").and_then(|v| v.as_str()) { + return format!("{tool_name}: {}", truncate_input_summary(command, 100)); + } + if let Some(path) = value.get("path").and_then(|v| v.as_str()) { + return format!("{tool_name}: {}", truncate_input_summary(path, 100)); + } + if let Some(file_path) = value.get("file_path").and_then(|v| v.as_str()) { + return format!("{tool_name}: {}", truncate_input_summary(file_path, 100)); + } + } + format!("{tool_name}({})", truncate_input_summary(input_json, 60)) +} + +fn truncate_input_summary(input: &str, max_bytes: usize) -> String { + if input.len() <= max_bytes { + input.to_string() + } else { + let end = (0..=max_bytes) + .rev() + .find(|&index| input.is_char_boundary(index)) + .unwrap_or_default(); + let mut truncated = input[..end].to_string(); + truncated.push_str("..."); + truncated + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + use crate::tool::ToolCall; + + fn tool_call(id: &str, name: &str) -> ToolCall { + ToolCall { + id: id.to_string(), + name: name.to_string(), + input: json!({}), + } + } + + fn tool_result(id: &str) -> AgentEvent { + AgentEvent::ToolExecutionFinished { + result: ContentBlock::ToolResult { + tool_use_id: id.to_string(), + content: mentra_provider::ToolResultContent::text("done"), + is_error: false, + }, + } + } + + #[test] + fn tool_completion_names_the_tool_that_was_queued() { + let mut seq = 0; + let mut names = ToolNameIndex::default(); + + map_agent_event( + &AgentEvent::ToolUseReady { + index: 0, + call: tool_call("tc-1", "files"), + }, + &mut seq, + &mut names, + ); + let mapped = map_agent_event(&tool_result("tc-1"), &mut seq, &mut names); + + assert!(matches!( + &mapped[0].1, + SessionEvent::ToolCompleted { tool_call_id, tool_name, .. } + if tool_call_id == "tc-1" && tool_name == "files" + )); + } + + #[test] + fn a_call_that_only_started_is_still_named_on_completion() { + let mut seq = 0; + let mut names = ToolNameIndex::default(); + + map_agent_event( + &AgentEvent::ToolExecutionStarted { + call: tool_call("tc-2", "shell"), + }, + &mut seq, + &mut names, + ); + let mapped = map_agent_event(&tool_result("tc-2"), &mut seq, &mut names); + + assert!(matches!( + &mapped[0].1, + SessionEvent::ToolCompleted { tool_name, .. } if tool_name == "shell" + )); + } + + #[test] + fn concurrent_calls_do_not_borrow_each_others_names() { + let mut seq = 0; + let mut names = ToolNameIndex::default(); + + for (id, name) in [("tc-a", "files"), ("tc-b", "shell")] { + map_agent_event( + &AgentEvent::ToolUseReady { + index: 0, + call: tool_call(id, name), + }, + &mut seq, + &mut names, + ); + } + + // Completions arrive in the opposite order to the queueing. + let second = map_agent_event(&tool_result("tc-b"), &mut seq, &mut names); + let first = map_agent_event(&tool_result("tc-a"), &mut seq, &mut names); + + assert!(matches!( + &second[0].1, + SessionEvent::ToolCompleted { tool_name, .. } if tool_name == "shell" + )); + assert!(matches!( + &first[0].1, + SessionEvent::ToolCompleted { tool_name, .. } if tool_name == "files" + )); + } + + #[test] + fn a_result_for_an_unseen_call_still_maps_with_an_empty_name() { + let mut seq = 0; + let mut names = ToolNameIndex::default(); + + let mapped = map_agent_event(&tool_result("never-queued"), &mut seq, &mut names); + + assert!(matches!( + &mapped[0].1, + SessionEvent::ToolCompleted { tool_call_id, tool_name, .. } + if tool_call_id == "never-queued" && tool_name.is_empty() + )); + } + + #[test] + fn a_completed_call_is_forgotten() { + let mut names = ToolNameIndex::default(); + names.remember("tc-1", "files"); + + assert_eq!(names.resolve("tc-1").as_deref(), Some("files")); + assert_eq!(names.resolve("tc-1"), None, "the entry must not linger"); + } + + #[test] + fn text_delta_maps_to_assistant_token_delta() { + let event = AgentEvent::TextDelta { + delta: "hi".to_string(), + full_text: "hi".to_string(), + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::AssistantTokenDelta { delta, .. } if delta == "hi" + )); + assert_eq!(seq, 1); + } + + #[test] + fn reasoning_delta_maps_to_assistant_reasoning_delta() { + let event = AgentEvent::ReasoningDelta { + delta: "private".to_string(), + full_text: "private chain".to_string(), + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::AssistantReasoningDelta { delta, full_text } + if delta == "private" && full_text == "private chain" + )); + assert_eq!(seq, 1); + } + + #[test] + fn tool_use_ready_maps_to_tool_queued() { + let event = AgentEvent::ToolUseReady { + index: 0, + call: ToolCall { + id: "tc-1".to_string(), + name: "read".to_string(), + input: json!({"path": "/foo"}), + }, + }; + let mut seq = 10; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::ToolQueued { tool_call_id, tool_name, .. } + if tool_call_id == "tc-1" && tool_name == "read" + )); + assert_eq!(mapped[0].0, 10); + assert_eq!(seq, 11); + } + + #[test] + fn tool_execution_finished_maps_to_tool_completed() { + let event = AgentEvent::ToolExecutionFinished { + result: ContentBlock::ToolResult { + tool_use_id: "tc-2".to_string(), + content: mentra_provider::ToolResultContent::text("ok"), + is_error: false, + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::ToolCompleted { tool_call_id, is_error, .. } + if tool_call_id == "tc-2" && !is_error + )); + } + + #[test] + fn compaction_maps_to_started_and_completed() { + let event = AgentEvent::ContextCompacted { + details: CompactionDetails { + trigger: crate::agent::CompactionTrigger::Auto, + mode: crate::compaction::CompactionExecutionMode::Local, + agent_id: "a1".to_string(), + transcript_path: std::path::PathBuf::from("/tmp"), + replaced_items: 10, + preserved_items: 5, + preserved_user_turns: 2, + preserved_delegation_results: 1, + resulting_transcript_len: 7, + extracted_facts_count: 0, + summary_preview: String::new(), + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 2); + assert!(matches!( + &mapped[0].1, + SessionEvent::CompactionStarted { .. } + )); + assert!(matches!( + &mapped[1].1, + SessionEvent::CompactionCompleted { .. } + )); + assert_eq!(seq, 2); + } + + #[test] + fn run_started_maps_to_empty() { + let event = AgentEvent::RunStarted; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert!(mapped.is_empty()); + assert_eq!(seq, 0); + } + + // --- classify_mutability tests --- + + #[test] + fn classify_mutability_no_side_effects_is_read_only() { + let result = classify_mutability( + ToolSideEffectLevel::None, + ToolExecutionCategory::ExclusiveLocalMutation, + ); + assert_eq!(result, ToolMutability::ReadOnly); + } + + #[test] + fn classify_mutability_read_only_parallel_is_read_only() { + let result = classify_mutability( + ToolSideEffectLevel::Process, + ToolExecutionCategory::ReadOnlyParallel, + ); + assert_eq!(result, ToolMutability::ReadOnly); + } + + #[test] + fn classify_mutability_side_effects_exclusive_is_mutating() { + let result = classify_mutability( + ToolSideEffectLevel::LocalState, + ToolExecutionCategory::ExclusiveLocalMutation, + ); + assert_eq!(result, ToolMutability::Mutating); + } + + #[test] + fn classify_mutability_external_delegation_is_mutating() { + let result = classify_mutability( + ToolSideEffectLevel::External, + ToolExecutionCategory::Delegation, + ); + assert_eq!(result, ToolMutability::Mutating); + } + + #[test] + fn classify_mutability_none_with_read_only_parallel_is_read_only() { + let result = classify_mutability( + ToolSideEffectLevel::None, + ToolExecutionCategory::ReadOnlyParallel, + ); + assert_eq!(result, ToolMutability::ReadOnly); + } + + // --- derive_tool_summary tests --- + + #[test] + fn derive_tool_summary_extracts_command() { + let summary = derive_tool_summary("shell", r#"{"command":"ls -la /tmp"}"#); + assert_eq!(summary, "shell: ls -la /tmp"); + } + + #[test] + fn derive_tool_summary_extracts_path() { + let summary = derive_tool_summary("read", r#"{"path":"/home/user/file.rs"}"#); + assert_eq!(summary, "read: /home/user/file.rs"); + } + + #[test] + fn derive_tool_summary_extracts_file_path() { + let summary = + derive_tool_summary("write", r#"{"file_path":"/tmp/out.txt","content":"hi"}"#); + assert_eq!(summary, "write: /tmp/out.txt"); + } + + #[test] + fn derive_tool_summary_falls_back_to_raw_input() { + let summary = derive_tool_summary("custom", r#"{"foo":"bar"}"#); + assert_eq!(summary, r#"custom({"foo":"bar"})"#); + } + + #[test] + fn derive_tool_summary_handles_invalid_json() { + let summary = derive_tool_summary("broken", "not json at all"); + assert_eq!(summary, "broken(not json at all)"); + } + + #[test] + fn derive_tool_summary_truncates_long_command() { + let long_cmd = "x".repeat(200); + let input = format!(r#"{{"command":"{long_cmd}"}}"#); + let summary = derive_tool_summary("shell", &input); + assert!(summary.len() < 200); + assert!(summary.ends_with("...")); + } + + #[test] + fn input_summary_truncates_before_a_utf8_boundary() { + let prefix = "x".repeat(199); + let input = format!("{prefix}—tail"); + + assert_eq!(truncate_input_summary(&input, 200), format!("{prefix}...")); + } + + // --- ToolExecutionProgress mapping test --- + + #[test] + fn tool_execution_progress_maps_to_tool_progress() { + let event = AgentEvent::ToolExecutionProgress { + id: "tc-5".to_string(), + name: "shell".to_string(), + progress: "50% complete".to_string(), + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::ToolProgress { tool_call_id, tool_name, progress } + if tool_call_id == "tc-5" && tool_name == "shell" && progress == "50% complete" + )); + assert_eq!(seq, 1); + } + + // --- tool_use_ready now uses derive_tool_summary --- + + #[test] + fn tool_use_ready_summary_uses_path_field() { + let event = AgentEvent::ToolUseReady { + index: 0, + call: ToolCall { + id: "tc-10".to_string(), + name: "read".to_string(), + input: json!({"path": "/src/main.rs"}), + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + if let SessionEvent::ToolQueued { summary, .. } = &mapped[0].1 { + assert_eq!(summary, "read: /src/main.rs"); + } else { + panic!("expected ToolQueued"); + } + } + + // --- SubagentSpawned / SubagentFinished mapping tests --- + + #[test] + fn subagent_spawned_maps_to_task_updated_spawned() { + let event = AgentEvent::SubagentSpawned { + agent: SpawnedAgentSummary { + id: "sub-1".to_string(), + name: "researcher".to_string(), + model: "mock-model".to_string(), + status: SpawnedAgentStatus::Running, + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + task_id, + kind: TaskKind::Subagent, + status: TaskLifecycleStatus::Spawned, + title, + detail: None, + } + if task_id == "sub-1" && title == "researcher" + )); + assert_eq!(seq, 1); + } + + #[test] + fn subagent_finished_success_maps_to_task_updated_finished() { + let event = AgentEvent::SubagentFinished { + agent: SpawnedAgentSummary { + id: "sub-2".to_string(), + name: "analyst".to_string(), + model: "mock-model".to_string(), + status: SpawnedAgentStatus::Finished, + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + task_id, + kind: TaskKind::Subagent, + status: TaskLifecycleStatus::Finished, + title, + detail: None, + } + if task_id == "sub-2" && title == "analyst" + )); + } + + #[test] + fn subagent_finished_failure_maps_to_task_updated_failed_with_detail() { + let event = AgentEvent::SubagentFinished { + agent: SpawnedAgentSummary { + id: "sub-3".to_string(), + name: "writer".to_string(), + model: "mock-model".to_string(), + status: SpawnedAgentStatus::Failed("provider timeout".to_string()), + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + task_id, + kind: TaskKind::Subagent, + status: TaskLifecycleStatus::Failed, + title, + detail: Some(msg), + } + if task_id == "sub-3" && title == "writer" && msg == "provider timeout" + )); + } + + // --- BackgroundTaskStarted / BackgroundTaskFinished mapping tests --- + + #[test] + fn background_task_started_maps_to_task_updated_running() { + let event = AgentEvent::BackgroundTaskStarted { + task: BackgroundTaskSummary { + id: "bg-1".to_string(), + command: "cargo test".to_string(), + cwd: std::path::PathBuf::from("/tmp"), + status: BackgroundTaskStatus::Running, + output_preview: None, + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + task_id, + kind: TaskKind::BackgroundTask, + status: TaskLifecycleStatus::Running, + title, + detail: None, + } + if task_id == "bg-1" && title == "cargo test" + )); + } + + #[test] + fn background_task_finished_success_maps_to_task_updated_finished() { + let event = AgentEvent::BackgroundTaskFinished { + task: BackgroundTaskSummary { + id: "bg-2".to_string(), + command: "npm run build".to_string(), + cwd: std::path::PathBuf::from("/project"), + status: BackgroundTaskStatus::Finished, + output_preview: Some("Build complete".to_string()), + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + task_id, + kind: TaskKind::BackgroundTask, + status: TaskLifecycleStatus::Finished, + title, + detail: Some(preview), + } + if task_id == "bg-2" && title == "npm run build" && preview == "Build complete" + )); + } + + #[test] + fn background_task_finished_failure_maps_to_task_updated_failed() { + let event = AgentEvent::BackgroundTaskFinished { + task: BackgroundTaskSummary { + id: "bg-3".to_string(), + command: "make".to_string(), + cwd: std::path::PathBuf::from("/build"), + status: BackgroundTaskStatus::Failed, + output_preview: Some("exit code 2".to_string()), + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + task_id, + kind: TaskKind::BackgroundTask, + status: TaskLifecycleStatus::Failed, + title, + detail: Some(preview), + } + if task_id == "bg-3" && title == "make" && preview == "exit code 2" + )); + } + + // --- TeammateSpawned / TeammateUpdated mapping tests --- + + #[test] + fn teammate_spawned_maps_to_task_updated_spawned() { + let event = AgentEvent::TeammateSpawned { + teammate: TeamMemberSummary { + id: "tm-1".to_string(), + name: "reviewer".to_string(), + role: "code review".to_string(), + model: "mock-model".to_string(), + status: TeamMemberStatus::Idle, + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + task_id, + kind: TaskKind::Teammate, + status: TaskLifecycleStatus::Spawned, + title, + detail: Some(role), + } + if task_id == "tm-1" && title == "reviewer" && role == "code review" + )); + } + + #[test] + fn teammate_updated_shutdown_maps_to_finished() { + let event = AgentEvent::TeammateUpdated { + teammate: TeamMemberSummary { + id: "tm-2".to_string(), + name: "tester".to_string(), + role: "testing".to_string(), + model: "mock-model".to_string(), + status: TeamMemberStatus::Shutdown, + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + status: TaskLifecycleStatus::Finished, + .. + } + )); + } + + #[test] + fn teammate_updated_failed_maps_to_failed_with_message() { + let event = AgentEvent::TeammateUpdated { + teammate: TeamMemberSummary { + id: "tm-3".to_string(), + name: "deployer".to_string(), + role: "deploy".to_string(), + model: "mock-model".to_string(), + status: TeamMemberStatus::Failed("connection refused".to_string()), + }, + }; + let mut seq = 0; + let mapped = map_agent_event(&event, &mut seq, &mut ToolNameIndex::default()); + assert_eq!(mapped.len(), 1); + assert!(matches!( + &mapped[0].1, + SessionEvent::TaskUpdated { + status: TaskLifecycleStatus::Failed, + detail: Some(msg), + .. + } + if msg == "connection refused" + )); + } +} diff --git a/vendor/mentra/src/session/permission.rs b/vendor/mentra/src/session/permission.rs new file mode 100644 index 0000000..dce54d7 --- /dev/null +++ b/vendor/mentra/src/session/permission.rs @@ -0,0 +1,878 @@ +mod pattern; + +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use tokio::sync::{broadcast, oneshot}; + +use super::event::{PermissionRuleScope, SessionEvent}; +use crate::{ + runtime::RuntimeError, + tool::{ + ToolAuthorizationDecision, ToolAuthorizationOutcome, ToolAuthorizationRequest, + ToolAuthorizer, + }, +}; + +/// A pending permission request awaiting a UI decision. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct PermissionRequest { + pub request_id: String, + pub tool_call_id: String, + pub tool_name: String, + pub description: String, + /// JSON-encoded preview data. Stored as `String` because + /// `serde_json::Value` does not implement `Eq`. + pub preview: String, +} + +/// What a refusal says when the deciding layer offered no reason of its own. +const DENIED_BY_SESSION_APPROVER: &str = "denied by session approver"; + +/// What a remembered refusal says when the rule it was stored as kept no reason. +const BLOCKED_BY_REMEMBERED_RULE: &str = "blocked by remembered session rule"; + +/// What the model reads when a remembered rule refuses a call. +/// +/// The words the host first refused with come back in front, because they are +/// the part that says what to do instead; the rest says the answer is standing, +/// because a model told only that something was blocked asks again, and asking +/// again is the one thing that cannot change a remembered rule. +fn remembered_denial(reason: Option<&str>) -> String { + match reason { + Some(reason) => format!( + "{reason} — remembered from an earlier refusal, so asking again will not change it" + ), + None => BLOCKED_BY_REMEMBERED_RULE.to_string(), + } +} + +/// The response to a permission request from the UI layer. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PermissionDecision { + pub allow: bool, + pub remember_as: Option, + /// Why the call was refused, in the words the model will read. + /// + /// A denial reaches the model as the tool's result, so what it says + /// changes what the model does next: told only that something was denied + /// it tries the write again, told that this run does not allow writes it + /// stops and reports. Set it with [`PermissionDecision::with_reason`]. + /// Ignored when `allow` is set, and a refusal that leaves it unset still + /// reads "denied by session approver" as it always has. + pub reason: Option, +} + +impl PermissionDecision { + /// Allow the tool call without remembering. + pub fn allow() -> Self { + Self { + allow: true, + remember_as: None, + reason: None, + } + } + + /// Deny the tool call without remembering. + pub fn deny() -> Self { + Self { + allow: false, + remember_as: None, + reason: None, + } + } + + /// Allow the tool call and remember the decision for the given scope. + pub fn allow_and_remember(scope: PermissionRuleScope) -> Self { + Self { + allow: true, + remember_as: Some(scope), + reason: None, + } + } + + /// Deny the tool call and remember the decision for the given scope. + pub fn deny_and_remember(scope: PermissionRuleScope) -> Self { + Self { + allow: false, + remember_as: Some(scope), + reason: None, + } + } + + /// The same decision, carrying the reason the model should read. + /// + /// Only refusals have anything to explain: an allowed call explains + /// itself by happening. + pub fn with_reason(self, reason: impl Into) -> Self { + Self { + reason: Some(reason.into()), + ..self + } + } +} + +/// Key for looking up remembered permission rules. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +pub struct RuleKey { + pub tool_name: String, + /// Wildcard pattern matched against the JSON encoding of the call's + /// structured input, or `None` to answer every call to the tool. + /// + /// Matched as data rather than as a path: `*` matches any run of + /// characters including `/`, `**` means the same as `*`, `?` matches one + /// character, and every other character — JSON's braces, brackets and + /// commas included — is literal. Matching is anchored, so a rule about a + /// fragment is written `*fragment*`. + /// + /// Path-glob semantics were wrong here: `*` stopped at `/`, so any preview + /// carrying an absolute path made every key serialized after it + /// unmatchable, and a rule written against one silently answered nothing. + pub pattern: Option, +} + +/// A stored permission rule that was previously decided by the user. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct RememberedRule { + pub key: RuleKey, + pub allow: bool, + pub scope: PermissionRuleScope, + /// Why the remembered refusal refused, in the words the model will read. + /// + /// A remembered rule answers every later call itself, without ever + /// reaching the approver again, so a rule that keeps the verdict and drops + /// the reason lets the host explain itself exactly once: every repeat after + /// that reads only that something was blocked. Written from + /// [`PermissionDecision::reason`] when the remembered decision is a + /// refusal, and left unset for an allow, which explains itself by + /// happening. A refusal that kept no reason still reads "blocked by + /// remembered session rule" as it always has. + /// + /// `serde(default)` keeps rules persisted before this field existed + /// deserializing unchanged. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reason: Option, +} + +/// Thread-safe in-memory store for remembered permission rules. +#[derive(Debug, Clone)] +pub struct RuleStore { + inner: Arc>>, +} + +impl Default for RuleStore { + fn default() -> Self { + Self::new() + } +} + +impl RuleStore { + /// Creates an empty rule store. + pub fn new() -> Self { + Self { + inner: Arc::new(Mutex::new(HashMap::new())), + } + } + + /// Adds or overwrites a remembered rule. + pub fn add_rule(&self, rule: RememberedRule) { + let mut rules = self.inner.lock().unwrap_or_else(|e| e.into_inner()); + rules.insert(rule.key.clone(), rule); + } + + /// Checks whether a tool is allowed by a remembered rule. + /// + /// Pattern rules are matched against `input_json` with the wildcard syntax + /// documented on [`RuleKey::pattern`] and take precedence over bare + /// (no-pattern) rules. Returns `Some(true)` if allowed, `Some(false)` if + /// denied, or `None` if no matching rule exists. + /// Use [`RuleStore::matching_rule`] when the rule's own reason matters. + pub fn check(&self, tool_name: &str, input_json: Option<&str>) -> Option { + self.matching_rule(tool_name, input_json) + .map(|rule| rule.allow) + } + + /// The remembered rule that answers a call, if one does. + /// + /// Matches exactly as [`RuleStore::check`] does — pattern rules against + /// `input_json` by wildcard, taking precedence over bare (no-pattern) + /// rules — + /// and hands back the whole rule, so a refusal can restate the reason it + /// was remembered with rather than only its verdict. + pub fn matching_rule( + &self, + tool_name: &str, + input_json: Option<&str>, + ) -> Option { + let rules = self.inner.lock().unwrap_or_else(|e| e.into_inner()); + + let mut pattern_match: Option<&RememberedRule> = None; + let mut bare_match: Option<&RememberedRule> = None; + + for rule in rules.values() { + if rule.key.tool_name != tool_name { + continue; + } + match &rule.key.pattern { + Some(rule_pattern) => { + if let Some(json) = input_json + && pattern::matches(rule_pattern, json) + { + pattern_match = Some(rule); + } + } + None => { + bare_match = Some(rule); + } + } + } + + pattern_match.or(bare_match).cloned() + } + + /// Returns all remembered rules as a vector. + pub fn rules(&self) -> Vec { + let rules = self.inner.lock().unwrap_or_else(|e| e.into_inner()); + rules.values().cloned().collect() + } + + /// Removes all rules that match the given scope. + pub fn clear_scope(&self, scope: PermissionRuleScope) { + let mut rules = self.inner.lock().unwrap_or_else(|e| e.into_inner()); + rules.retain(|_, rule| rule.scope != scope); + } +} + +/// Thread-safe store for pending permission requests that can be resolved later. +#[derive(Debug, Clone, Default)] +pub(crate) struct PendingPermissionStore { + inner: Arc>>, +} + +impl PendingPermissionStore { + pub(crate) fn new() -> Self { + Self::default() + } + + pub(crate) fn insert(&self, request_id: String, entry: PendingPermissionEntry) { + let mut pending = self.inner.lock().unwrap_or_else(|e| e.into_inner()); + pending.insert(request_id, entry); + } + + pub(crate) fn remove(&self, request_id: &str) -> Option { + let mut pending = self.inner.lock().unwrap_or_else(|e| e.into_inner()); + pending.remove(request_id) + } + + #[cfg(test)] + pub(crate) fn contains(&self, request_id: &str) -> bool { + let pending = self.inner.lock().unwrap_or_else(|e| e.into_inner()); + pending.contains_key(request_id) + } +} + +/// Internal entry tracking a pending permission with its oneshot response channel. +#[derive(Debug)] +pub(crate) struct PendingPermissionEntry { + pub(crate) tool_call_id: String, + pub(crate) tool_name: String, + pub(crate) sender: oneshot::Sender, +} + +/// Session-scoped wrapper around the runtime tool authorizer. +/// +/// This is the bridge that turns `Prompt` outcomes into typed +/// `SessionEvent::PermissionRequested` events, stores the pending request, and +/// suspends execution until a matching decision arrives. +#[derive(Clone)] +pub(crate) struct SessionToolAuthorizer { + inner: Option>, + event_tx: broadcast::Sender, + pending_permissions: PendingPermissionStore, + rule_store: RuleStore, +} + +impl SessionToolAuthorizer { + pub(crate) fn new( + inner: Option>, + event_tx: broadcast::Sender, + pending_permissions: PendingPermissionStore, + rule_store: RuleStore, + ) -> Self { + Self { + inner, + event_tx, + pending_permissions, + rule_store, + } + } +} + +#[async_trait] +impl ToolAuthorizer for SessionToolAuthorizer { + async fn authorize( + &self, + request: &ToolAuthorizationRequest, + ) -> Result { + let input_json = serde_json::to_string(&request.preview.structured_input).ok(); + if let Some(rule) = self + .rule_store + .matching_rule(&request.tool_name, input_json.as_deref()) + { + return Ok(if rule.allow { + ToolAuthorizationDecision::allow() + } else { + // The approver is not consulted again, so the rule is the only + // place the original reason can still come from. + ToolAuthorizationDecision::deny(remembered_denial(rule.reason.as_deref())) + }); + } + + let Some(inner) = &self.inner else { + return Ok(ToolAuthorizationDecision::allow()); + }; + + let decision = inner.authorize(request).await?; + if decision.outcome != ToolAuthorizationOutcome::Prompt { + return Ok(decision); + } + + let request_id = format!("perm-{}", request.tool_call_id); + let description = decision + .reason + .clone() + .unwrap_or_else(|| format!("Approval required for {}", request.tool_name)); + let preview = serde_json::to_string(&request.preview.structured_input) + .unwrap_or_else(|_| "{}".to_string()); + let (sender, receiver) = oneshot::channel(); + + self.pending_permissions.insert( + request_id.clone(), + PendingPermissionEntry { + tool_call_id: request.tool_call_id.clone(), + tool_name: request.tool_name.clone(), + sender, + }, + ); + + let _ = self.event_tx.send(SessionEvent::PermissionRequested { + request_id: request_id.clone(), + tool_call_id: request.tool_call_id.clone(), + tool_name: request.tool_name.clone(), + description, + preview, + }); + + let resolved = receiver + .await + .unwrap_or_else(|_| PermissionDecision::deny()); + Ok(if resolved.allow { + ToolAuthorizationDecision::allow() + } else { + // Whoever answered gets to say why, because that text is what the + // model reads; a refusal that explains nothing keeps the wording + // this has always used. + ToolAuthorizationDecision::deny( + resolved + .reason + .unwrap_or_else(|| DENIED_BY_SESSION_APPROVER.to_string()), + ) + }) + } + + fn timeout(&self) -> Option { + self.inner + .as_ref() + .and_then(|authorizer| authorizer.timeout()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + use crate::tool::{ + ToolApprovalCategory, ToolAuthorizationPreview, ToolCapability, ToolDurability, + ToolExecutionCategory, ToolSideEffectLevel, + }; + + #[derive(Clone)] + struct PromptAuthorizer; + + #[async_trait] + impl ToolAuthorizer for PromptAuthorizer { + async fn authorize( + &self, + _request: &ToolAuthorizationRequest, + ) -> Result { + Ok(ToolAuthorizationDecision::prompt("needs manual review")) + } + } + + /// Refuses in words no remembered rule would use, so a call that reached + /// the approver shows up as a wrong string rather than as a test that + /// blocks forever waiting for an answer nobody will give. + #[derive(Clone)] + struct ApproverOfLastResort; + + #[async_trait] + impl ToolAuthorizer for ApproverOfLastResort { + async fn authorize( + &self, + _request: &ToolAuthorizationRequest, + ) -> Result { + Ok(ToolAuthorizationDecision::deny( + "the approver was asked again", + )) + } + } + + fn sample_request() -> ToolAuthorizationRequest { + ToolAuthorizationRequest { + agent_id: "agent-1".to_string(), + agent_name: "agent".to_string(), + model: "mock-model".to_string(), + history_len: 3, + tool_call_id: "tool-1".to_string(), + tool_name: "shell".to_string(), + preview: ToolAuthorizationPreview { + working_directory: std::env::temp_dir(), + capabilities: vec![ToolCapability::ProcessExec], + side_effect_level: ToolSideEffectLevel::Process, + durability: ToolDurability::Ephemeral, + execution_category: ToolExecutionCategory::ExclusiveLocalMutation, + approval_category: ToolApprovalCategory::Process, + raw_input: json!({ "command": "cargo test" }), + structured_input: json!({ "kind": "shell", "command": "cargo test" }), + }, + } + } + + #[tokio::test] + async fn session_tool_authorizer_emits_permission_request_and_waits() { + let (event_tx, mut rx) = broadcast::channel(8); + let pending = PendingPermissionStore::new(); + let authorizer = SessionToolAuthorizer::new( + Some(Arc::new(PromptAuthorizer)), + event_tx, + pending.clone(), + RuleStore::new(), + ); + let request = sample_request(); + + let authorize_task = tokio::spawn({ + let authorizer = authorizer.clone(); + let request = request.clone(); + async move { authorizer.authorize(&request).await.unwrap() } + }); + + let event = tokio::time::timeout(Duration::from_millis(200), rx.recv()) + .await + .expect("permission request should arrive") + .expect("event should be present"); + + let request_id = match event { + SessionEvent::PermissionRequested { + request_id, + tool_call_id, + tool_name, + .. + } => { + assert_eq!(tool_call_id, "tool-1"); + assert_eq!(tool_name, "shell"); + request_id + } + other => panic!("expected PermissionRequested, got {other:?}"), + }; + + assert!(pending.contains(&request_id)); + let entry = pending + .remove(&request_id) + .expect("pending permission should be registered"); + entry + .sender + .send(PermissionDecision::allow()) + .expect("decision send should succeed"); + + let decision = tokio::time::timeout(Duration::from_millis(200), authorize_task) + .await + .expect("authorization should resume") + .expect("task should succeed"); + assert_eq!(decision.outcome, ToolAuthorizationOutcome::Allow); + } + + /// Runs one authorize-and-resolve round trip, answering with `decision`, + /// and returns what the authorizer handed back to the tool loop. + async fn resolved_with(decision: PermissionDecision) -> ToolAuthorizationDecision { + let (event_tx, mut rx) = broadcast::channel(8); + let pending = PendingPermissionStore::new(); + let authorizer = SessionToolAuthorizer::new( + Some(Arc::new(PromptAuthorizer)), + event_tx, + pending.clone(), + RuleStore::new(), + ); + + let authorize_task = tokio::spawn({ + let authorizer = authorizer.clone(); + async move { authorizer.authorize(&sample_request()).await.unwrap() } + }); + + let event = tokio::time::timeout(Duration::from_millis(200), rx.recv()) + .await + .expect("permission request should arrive") + .expect("event should be present"); + let SessionEvent::PermissionRequested { request_id, .. } = &event else { + panic!("expected PermissionRequested, got {event:?}"); + }; + + pending + .remove(request_id) + .expect("pending permission should be registered") + .sender + .send(decision) + .expect("decision send should succeed"); + + tokio::time::timeout(Duration::from_millis(200), authorize_task) + .await + .expect("authorization should resume") + .expect("task should succeed") + } + + #[tokio::test] + async fn a_reasoned_denial_carries_its_words_to_the_tool_result() { + // The reason becomes the tool result the model reads, so anything + // rewritten or dropped on the way is a reason it never sees. + let decision = + resolved_with(PermissionDecision::deny().with_reason("this run does not allow writes")) + .await; + + assert_eq!(decision.outcome, ToolAuthorizationOutcome::Deny); + assert_eq!( + decision.reason.as_deref(), + Some("this run does not allow writes") + ); + } + + #[tokio::test] + async fn a_denial_with_nothing_to_say_keeps_the_standing_wording() { + let decision = resolved_with(PermissionDecision::deny()).await; + + assert_eq!(decision.outcome, ToolAuthorizationOutcome::Deny); + assert_eq!(decision.reason.as_deref(), Some(DENIED_BY_SESSION_APPROVER)); + } + + #[tokio::test] + async fn a_reason_on_an_allowed_call_changes_nothing() { + let decision = resolved_with(PermissionDecision::allow().with_reason("ignored")).await; + + assert_eq!(decision.outcome, ToolAuthorizationOutcome::Allow); + assert_eq!( + decision.reason, None, + "an allowed call has nothing to explain" + ); + } + + /// Answers one authorize call from `store` alone. The approver behind it + /// refuses in its own words, so a rule that failed to answer is visible. + async fn answered_by_rule(store: RuleStore) -> ToolAuthorizationDecision { + let (event_tx, _rx) = broadcast::channel(8); + let authorizer = SessionToolAuthorizer::new( + Some(Arc::new(ApproverOfLastResort)), + event_tx, + PendingPermissionStore::new(), + store, + ); + + authorizer + .authorize(&sample_request()) + .await + .expect("authorization should resolve") + } + + /// A bare `shell` rule for the session, remembered with `reason` or without. + fn shell_rule(allow: bool, reason: Option<&str>) -> RememberedRule { + RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow, + scope: PermissionRuleScope::Session, + reason: reason.map(str::to_owned), + } + } + + #[tokio::test] + async fn a_remembered_refusal_restates_the_reason_it_was_remembered_with() { + // Nothing asks the approver a second time, so the rule is the only + // thing left that knows why the first answer was no. + let store = RuleStore::new(); + store.add_rule(shell_rule(false, Some("this run does not allow writes"))); + + let decision = answered_by_rule(store).await; + + assert_eq!(decision.outcome, ToolAuthorizationOutcome::Deny); + assert_eq!( + decision.reason.as_deref(), + Some( + "this run does not allow writes — remembered from an earlier refusal, so asking again will not change it" + ) + ); + } + + #[tokio::test] + async fn a_refusal_remembered_without_a_reason_keeps_the_standing_wording() { + let store = RuleStore::new(); + store.add_rule(shell_rule(false, None)); + + let decision = answered_by_rule(store).await; + + assert_eq!(decision.outcome, ToolAuthorizationOutcome::Deny); + assert_eq!(decision.reason.as_deref(), Some(BLOCKED_BY_REMEMBERED_RULE)); + } + + #[tokio::test] + async fn a_remembered_allow_answers_without_words() { + // Nothing writes a reason onto an allow, but the type permits one, and + // an allowed call still explains itself by happening. + let store = RuleStore::new(); + store.add_rule(shell_rule(true, Some("should never be read"))); + + let decision = answered_by_rule(store).await; + + assert_eq!(decision.outcome, ToolAuthorizationOutcome::Allow); + assert_eq!(decision.reason, None); + } + + #[test] + fn matching_rule_hands_back_the_reason_of_the_rule_that_won() { + // Precedence decides which reason the model reads, so the pattern + // rule's words must come back rather than the bare rule's. + let store = RuleStore::new(); + store.add_rule(shell_rule(false, Some("shell is refused in this run"))); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: Some("**cargo test**".to_owned()), + }, + allow: false, + scope: PermissionRuleScope::Session, + reason: Some("the test suite is not run from inside a run".to_owned()), + }); + + let matched = store + .matching_rule("shell", Some(r#"{"command":"cargo test"}"#)) + .expect("a rule should match"); + + assert_eq!( + matched.reason.as_deref(), + Some("the test suite is not run from inside a run") + ); + } + + #[test] + fn check_matches_tool_name_without_pattern() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + // Bare rule (no pattern) matches regardless of input_json content. + assert_eq!( + store.check("shell", Some(r#"{"command":"ls"}"#)), + Some(true) + ); + assert_eq!(store.check("shell", None), Some(true)); + } + + #[test] + fn check_matches_pattern_against_input_json() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: Some("*cargo test*".to_owned()), + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + assert_eq!( + store.check("shell", Some(r#"{"command":"cargo test"}"#)), + Some(true) + ); + } + + #[test] + fn check_pattern_rule_does_not_match_without_input() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: Some("*cargo test*".to_owned()), + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + // Pattern rule is ignored when input is None — no bare rule either, + // so result must be None. + assert_eq!(store.check("shell", None), None); + } + + #[test] + fn check_pattern_rule_takes_precedence_over_no_pattern() { + let store = RuleStore::new(); + // Bare rule: allow. + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + // Pattern rule: deny when input matches. `**` reads the same as `*` + // now that a pattern is matched as data, and is kept here because a + // rule persisted with that spelling has to keep answering. + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: Some("**rm -rf**".to_owned()), + }, + allow: false, + scope: PermissionRuleScope::Session, + reason: None, + }); + // Pattern match should win over the bare allow. + assert_eq!( + store.check("shell", Some(r#"{"command":"rm -rf /tmp"}"#)), + Some(false) + ); + } + + /// The preview a host builds for a routed command, with its keys in the + /// order `serde_json` writes them: an absolute `cwd` sits before `mode` + /// and `target`. + fn spawn_preview() -> &'static str { + r#"{"body":"cargo test","cwd":"/Users/dev/basis","mode":"command","target":"mac"}"# + } + + /// A pattern is matched against JSON, and JSON is not a path. Matched by a + /// path globber, `*` stops dead at the `/` inside an absolute `cwd`, so + /// every key serialized after `cwd` becomes unreachable — the rule saves, + /// reports nothing, and silently answers no call it was written for. + #[test] + fn a_pattern_reaches_a_key_that_follows_an_absolute_path() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "spawn".to_owned(), + pattern: Some(r#"**"mode":"command"**"#.to_owned()), + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + + assert_eq!(store.check("spawn", Some(spawn_preview())), Some(true)); + } + + #[test] + fn a_pattern_reaches_the_last_key_of_a_preview() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "spawn".to_owned(), + pattern: Some(r#"**"target":"mac"**"#.to_owned()), + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + + assert_eq!(store.check("spawn", Some(spawn_preview())), Some(true)); + } + + /// `**` was only ever needed because `*` could not cross a separator. + /// Both now mean the same thing, so a rule written either way answers. + #[test] + fn one_star_and_two_stars_both_cross_a_path_separator() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "spawn".to_owned(), + pattern: Some(r#"*"target":"mac"*"#.to_owned()), + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + + assert_eq!(store.check("spawn", Some(spawn_preview())), Some(true)); + } + + /// JSON is punctuation-dense, and a path globber reads some of that + /// punctuation as syntax: `{`…`}` is brace alternation and `[`…`]` a + /// character class. A pattern that quotes the front of an object must + /// match the object it quotes. + #[test] + fn json_punctuation_in_a_pattern_is_matched_literally() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "spawn".to_owned(), + pattern: Some(r#"{"body":"cargo test"*"#.to_owned()), + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + + assert_eq!(store.check("spawn", Some(spawn_preview())), Some(true)); + } + + #[test] + fn a_pattern_that_names_another_target_does_not_match() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "spawn".to_owned(), + pattern: Some(r#"**"target":"linux"**"#.to_owned()), + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + + assert_eq!(store.check("spawn", Some(spawn_preview())), None); + } + + #[test] + fn check_non_matching_pattern_falls_through() { + let store = RuleStore::new(); + // Only a pattern rule is present; input does not match it. + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: Some("*cargo test*".to_owned()), + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + // Non-matching input yields None (no bare fallback). + assert_eq!(store.check("shell", Some(r#"{"command":"ls"}"#)), None); + } +} diff --git a/vendor/mentra/src/session/permission/pattern.rs b/vendor/mentra/src/session/permission/pattern.rs new file mode 100644 index 0000000..a2dc57b --- /dev/null +++ b/vendor/mentra/src/session/permission/pattern.rs @@ -0,0 +1,160 @@ +//! Wildcard matching for remembered permission rule patterns. +//! +//! A rule pattern is matched against the JSON encoding of a tool call's +//! structured input. That string is data, not a filesystem path, and the +//! distinction is not cosmetic: a path globber stops `*` at `/`, so the moment +//! a preview carries an absolute path — a `cwd`, a file argument — every key +//! `serde_json` writes after it becomes unreachable. The rule still saves, the +//! store still reports success, and the operator is told nothing; the rule +//! simply never answers a call again. Path globbers also read JSON's own +//! punctuation as syntax, taking `{`…`}` for alternation and `[`…`]` for a +//! character class, so a pattern quoting the front of an object matches +//! something other than the object it quotes. +//! +//! So the syntax here is deliberately small and has no separator at all: +//! +//! - `*` matches any run of characters, `/` included, and `**` means the same +//! thing (patterns written when `**` was the only way to cross a separator +//! keep working unchanged). +//! - `?` matches exactly one character — one `char`, not one byte, so a +//! pattern stays predictable over non-ASCII input. +//! - Everything else, punctuation included, is literal. +//! +//! Matching is anchored: the pattern must account for the whole string, which +//! is why a substring rule is written `*needle*`. + +use std::str::Chars; + +/// Whether `pattern` matches the whole of `text`. +/// +/// Greedy with backtracking: each `*` first takes as little as possible and is +/// widened one character at a time only when the rest of the pattern fails, so +/// a match is found whenever one exists. +pub(crate) fn matches(pattern: &str, text: &str) -> bool { + let mut pattern_rest = pattern.chars(); + let mut text_rest = text.chars(); + // The pattern following the most recent `*`, paired with the text that + // star has not yet swallowed. Backtracking is letting that star eat one + // more character and retrying the rest of the pattern from there. Only the + // last star needs remembering: once a prefix has matched, an earlier star + // widening cannot rescue a failure a later one can. + let mut widest_star: Option<(Chars<'_>, Chars<'_>)> = None; + + loop { + let mut pattern_next = pattern_rest.clone(); + match pattern_next.next() { + Some('*') => { + pattern_rest = pattern_next; + widest_star = Some((pattern_rest.clone(), text_rest.clone())); + continue; + } + Some(expected) => { + let mut text_next = text_rest.clone(); + if let Some(actual) = text_next.next() + && (expected == '?' || expected == actual) + { + pattern_rest = pattern_next; + text_rest = text_next; + continue; + } + } + None => { + if text_rest.clone().next().is_none() { + return true; + } + } + } + + // The pattern cannot account for what is at this position. Widen the + // last star by one character, or admit there is nothing left to widen. + let Some((star_pattern, star_text)) = widest_star.as_mut() else { + return false; + }; + if star_text.next().is_none() { + return false; + } + pattern_rest = star_pattern.clone(); + text_rest = star_text.clone(); + } +} + +#[cfg(test)] +mod tests { + use super::matches; + + const PREVIEW: &str = + r#"{"body":"cargo test","cwd":"/Users/dev/basis","mode":"command","target":"mac"}"#; + + #[test] + fn a_star_crosses_a_path_separator() { + assert!(matches(r#"*"target":"mac"*"#, PREVIEW)); + assert!(matches(r#"*"mode":"command"*"#, PREVIEW)); + } + + #[test] + fn two_stars_mean_what_one_star_means() { + assert!(matches(r#"**"target":"mac"**"#, PREVIEW)); + assert!(matches("**", PREVIEW)); + assert!(matches("***", PREVIEW)); + } + + #[test] + fn json_punctuation_is_literal() { + assert!(matches(r#"{"body":"cargo test"*"#, PREVIEW)); + assert!(matches(r#"*"cwd":"/Users/dev/basis"*"#, PREVIEW)); + // Brace alternation and character classes are not syntax here, so a + // pattern naming them matches only text that contains them. + assert!(!matches("{a,b}", "a")); + assert!(matches("{a,b}", "{a,b}")); + assert!(!matches("[abc]", "a")); + assert!(matches("[abc]", "[abc]")); + } + + #[test] + fn a_question_mark_matches_exactly_one_character() { + assert!(matches("a?c", "abc")); + assert!(matches("a?c", "a/c")); + assert!(!matches("a?c", "ac")); + assert!(!matches("a?c", "abbc")); + } + + #[test] + fn a_question_mark_counts_characters_not_bytes() { + // One multi-byte char is one `?`, so a pattern behaves the same over + // text a host did not write in ASCII. + assert!(matches("a?c", "aéc")); + assert!(matches("?", "é")); + assert!(!matches("??", "é")); + } + + #[test] + fn matching_is_anchored() { + assert!(matches("cargo test", "cargo test")); + assert!(!matches("cargo", "cargo test")); + assert!(!matches("test", "cargo test")); + assert!(matches("cargo*", "cargo test")); + assert!(matches("*test", "cargo test")); + } + + #[test] + fn empty_pattern_matches_only_empty_text() { + assert!(matches("", "")); + assert!(!matches("", "a")); + assert!(matches("*", "")); + } + + #[test] + fn a_star_backtracks_until_the_rest_of_the_pattern_fits() { + // The first candidate for each star is wrong here; only widening in + // turn finds the match. + assert!(matches("*a*b", "xaybzb")); + assert!(matches("*ab*cd*", "zzabzzcdzz")); + assert!(!matches("*a*b", "xaybz")); + } + + #[test] + fn a_pattern_that_names_something_absent_does_not_match() { + assert!(!matches(r#"**"target":"linux"**"#, PREVIEW)); + assert!(!matches(r#"**"mode":"shell"**"#, PREVIEW)); + } +} diff --git a/vendor/mentra/src/session/tests.rs b/vendor/mentra/src/session/tests.rs new file mode 100644 index 0000000..dcae76f --- /dev/null +++ b/vendor/mentra/src/session/tests.rs @@ -0,0 +1,2485 @@ +#![expect( + clippy::unwrap_used, + reason = "tests use unwrap to fail immediately on invalid fixtures" +)] + +mod terminal_output; + +use crate::session::event::*; +use crate::session::permission::*; +use crate::session::types::*; + +// ---- Task 1 type-level tests (preserved) ---- + +#[test] +fn session_id_roundtrips_through_serde() { + let id = SessionId::new(); + let json = serde_json::to_string(&id).unwrap(); + let deserialized: SessionId = serde_json::from_str(&json).unwrap(); + assert_eq!(id, deserialized); +} + +#[test] +fn session_id_from_raw_preserves_value() { + let id = SessionId::from_raw("session-abc-123"); + assert_eq!(id.as_str(), "session-abc-123"); +} + +#[test] +fn session_metadata_serialization_roundtrip() { + let metadata = SessionMetadata::new( + SessionId::from_raw("session-test-1"), + "Test Session", + "claude-opus-4-20250514", + ); + let json = serde_json::to_value(&metadata).unwrap(); + let deserialized: SessionMetadata = serde_json::from_value(json).unwrap(); + assert_eq!(metadata, deserialized); +} + +#[test] +fn session_event_assistant_token_delta_roundtrip() { + let event = SessionEvent::AssistantTokenDelta { + delta: "hello".to_string(), + full_text: "hello".to_string(), + }; + let json = serde_json::to_value(&event).unwrap(); + assert_eq!(json["type"], "assistant_token_delta"); + let deserialized: SessionEvent = serde_json::from_value(json).unwrap(); + assert_eq!(event, deserialized); +} + +#[test] +fn session_event_assistant_reasoning_delta_roundtrip() { + let event = SessionEvent::AssistantReasoningDelta { + delta: "private".to_string(), + full_text: "private chain".to_string(), + }; + let json = serde_json::to_value(&event).unwrap(); + assert_eq!(json["type"], "assistant_reasoning_delta"); + let deserialized: SessionEvent = serde_json::from_value(json).unwrap(); + assert_eq!(event, deserialized); +} + +#[test] +fn session_event_tool_queued_roundtrip() { + let event = SessionEvent::ToolQueued { + tool_call_id: "tc-1".to_string(), + tool_name: "shell".to_string(), + summary: "Run 'cargo test'".to_string(), + mutability: ToolMutability::Mutating, + input_json: r#"{"command":"cargo test"}"#.to_string(), + }; + let json = serde_json::to_value(&event).unwrap(); + assert_eq!(json["type"], "tool_queued"); + assert_eq!(json["tool_name"], "shell"); + let deserialized: SessionEvent = serde_json::from_value(json).unwrap(); + assert_eq!(event, deserialized); +} + +#[test] +fn session_event_permission_requested_roundtrip() { + let preview_json = serde_json::to_string(&serde_json::json!({ + "command": "rm -rf /tmp/foo", + "cwd": "/Users/dev/project" + })) + .unwrap(); + let event = SessionEvent::PermissionRequested { + request_id: "perm-1".to_string(), + tool_call_id: "tc-1".to_string(), + tool_name: "shell".to_string(), + description: "Execute shell command: rm -rf /tmp/foo".to_string(), + preview: preview_json, + }; + let json = serde_json::to_value(&event).unwrap(); + assert_eq!(json["type"], "permission_requested"); + let deserialized: SessionEvent = serde_json::from_value(json).unwrap(); + assert_eq!(event, deserialized); +} + +#[test] +fn session_event_compaction_completed_roundtrip() { + let event = SessionEvent::CompactionCompleted { + agent_id: "agent-1".to_string(), + replaced_items: 42, + preserved_items: 8, + resulting_transcript_len: 10, + extracted_facts_count: 3, + summary_preview: "key facts extracted".to_string(), + }; + let json = serde_json::to_value(&event).unwrap(); + assert_eq!(json["type"], "compaction_completed"); + let deserialized: SessionEvent = serde_json::from_value(json).unwrap(); + assert_eq!(event, deserialized); +} + +#[test] +fn session_event_task_updated_roundtrip() { + let event = SessionEvent::TaskUpdated { + task_id: "bg-1".to_string(), + kind: TaskKind::BackgroundTask, + status: TaskLifecycleStatus::Running, + title: "cargo test -p mentra".to_string(), + detail: Some("exit code: 0".to_string()), + }; + let json = serde_json::to_value(&event).unwrap(); + assert_eq!(json["type"], "task_updated"); + let deserialized: SessionEvent = serde_json::from_value(json).unwrap(); + assert_eq!(event, deserialized); +} + +#[test] +fn all_session_event_variants_serialize_with_type_tag() { + let events: Vec = vec![ + SessionEvent::SessionStarted { + session_id: SessionId::from_raw("s1"), + }, + SessionEvent::UserMessage { + text: "hi".to_string(), + }, + SessionEvent::AssistantTokenDelta { + delta: "h".to_string(), + full_text: "h".to_string(), + }, + SessionEvent::AssistantReasoningDelta { + delta: "r".to_string(), + full_text: "r".to_string(), + }, + SessionEvent::AssistantMessageCompleted { + text: "hello".to_string(), + }, + SessionEvent::ToolQueued { + tool_call_id: "tc1".to_string(), + tool_name: "read".to_string(), + summary: "Read file".to_string(), + mutability: ToolMutability::ReadOnly, + input_json: "{}".to_string(), + }, + SessionEvent::ToolStarted { + tool_call_id: "tc1".to_string(), + tool_name: "read".to_string(), + }, + SessionEvent::ToolProgress { + tool_call_id: "tc1".to_string(), + tool_name: "read".to_string(), + progress: "50%".to_string(), + }, + SessionEvent::ToolCompleted { + tool_call_id: "tc1".to_string(), + tool_name: "read".to_string(), + summary: "Read 42 lines".to_string(), + is_error: false, + }, + SessionEvent::PermissionRequested { + request_id: "p1".to_string(), + tool_call_id: "tc1".to_string(), + tool_name: "shell".to_string(), + description: "run command".to_string(), + preview: "{}".to_string(), + }, + SessionEvent::PermissionResolved { + request_id: "p1".to_string(), + tool_call_id: "tc1".to_string(), + tool_name: "shell".to_string(), + outcome: PermissionOutcome::Allowed, + rule_scope: Some(PermissionRuleScope::Session), + }, + SessionEvent::TaskUpdated { + task_id: "t1".to_string(), + kind: TaskKind::Subagent, + status: TaskLifecycleStatus::Spawned, + title: "research".to_string(), + detail: None, + }, + SessionEvent::CompactionStarted { + agent_id: "a1".to_string(), + }, + SessionEvent::CompactionCompleted { + agent_id: "a1".to_string(), + replaced_items: 10, + preserved_items: 5, + resulting_transcript_len: 7, + extracted_facts_count: 0, + summary_preview: String::new(), + }, + SessionEvent::MemoryUpdated { + agent_id: "a1".to_string(), + stored_records: 3, + }, + SessionEvent::Notice { + severity: NoticeSeverity::Info, + message: "Context window 80% full".to_string(), + }, + SessionEvent::RetryAttempt { + agent_id: "a1".to_string(), + error_message: "transient error".to_string(), + attempt: 1, + max_attempts: 3, + next_delay_ms: 500, + }, + SessionEvent::Error { + message: "Provider timeout".to_string(), + recoverable: true, + }, + ]; + + for event in events { + let json = serde_json::to_value(&event).unwrap(); + assert!( + json.get("type").is_some(), + "Event missing 'type' tag: {event:?}" + ); + let roundtripped: SessionEvent = serde_json::from_value(json).unwrap(); + assert_eq!(event, roundtripped); + } +} + +// ---- Task 2 lifecycle tests ---- + +use crate::{ + ContentBlock, + runtime::{AgentStore, CancellationToken, RunOptions, SqliteRuntimeStore}, + test::MockRuntime, +}; + +#[tokio::test] +async fn create_session_produces_valid_metadata() { + let mock = MockRuntime::builder().text("hello").build().unwrap(); + let session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + assert_eq!(session.name(), "test-session"); + assert_eq!(session.metadata().title, "test-session"); + assert_eq!(session.metadata().model, mock.model().id); + assert_eq!(session.metadata().status, SessionStatus::Created); + assert_eq!(session.metadata().turn_count, 0); +} + +#[tokio::test] +async fn append_turn_returns_assistant_message() { + let mock = MockRuntime::builder() + .text("hello from session") + .build() + .unwrap(); + let mut session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + let message = session + .append_turn(vec![ContentBlock::text("hi")]) + .await + .unwrap(); + + assert_eq!(message.text(), "hello from session"); + assert_eq!(session.metadata().turn_count, 1); + assert_eq!(session.metadata().status, SessionStatus::Idle); +} + +#[tokio::test] +async fn append_turn_emits_user_and_assistant_events() { + let mock = MockRuntime::builder().text("response").build().unwrap(); + let mut session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + let mut rx = session.subscribe(); + + let _message = session + .append_turn(vec![ContentBlock::text("hello")]) + .await + .unwrap(); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + let has_user = events + .iter() + .any(|e| matches!(e, SessionEvent::UserMessage { text } if text == "hello")); + let has_assistant = events.iter().any( + |e| matches!(e, SessionEvent::AssistantMessageCompleted { text } if text == "response"), + ); + + assert!(has_user, "Expected UserMessage event, got: {events:?}"); + assert!( + has_assistant, + "Expected AssistantMessageCompleted event, got: {events:?}" + ); +} + +#[tokio::test] +async fn replay_returns_transcript_after_turn() { + let mock = MockRuntime::builder().text("world").build().unwrap(); + let mut session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + let _message = session + .append_turn(vec![ContentBlock::text("hello")]) + .await + .unwrap(); + + let transcript = session.replay(); + assert!( + !transcript.items().is_empty(), + "Transcript should have items after a turn" + ); +} + +#[tokio::test] +async fn a_cancelled_session_turn_fails_instead_of_running() { + let mock = MockRuntime::builder() + .text("never reached") + .build() + .unwrap(); + let mut session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + // Tripped before the turn starts, so the outcome does not depend on + // winning a race with the provider. + let cancellation = CancellationToken::default(); + cancellation.cancel(); + + let result = session + .append_turn_with_options( + vec![ContentBlock::text("go")], + RunOptions { + cancellation: Some(cancellation), + ..RunOptions::default() + }, + ) + .await; + + assert!( + result.is_err(), + "a cancelled turn must fail rather than run to completion" + ); + assert!( + matches!(session.metadata().status, SessionStatus::Failed(_)), + "the session must report the cancelled turn, not sit in Active" + ); +} + +#[tokio::test] +async fn a_session_turn_honors_a_token_budget() { + let mock = MockRuntime::builder() + .text("first") + .text("second") + .build() + .unwrap(); + let mut session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + // A budget of 1 is spent by the first round's reported usage, so the run + // stops gracefully at the next boundary rather than continuing. Pinned + // because `Session` passes `RunOptions` through and nothing else proves + // the budget survives that hop. + let message = session + .append_turn_with_options( + vec![ContentBlock::text("go")], + RunOptions { + token_budget: Some(1), + ..RunOptions::default() + }, + ) + .await + .unwrap(); + + assert_eq!(message.text(), "first"); + assert_eq!(session.metadata().status, SessionStatus::Idle); +} + +#[tokio::test] +async fn run_options_default_to_the_same_turn_append_turn_runs() { + let mock = MockRuntime::builder().text("hello").build().unwrap(); + let mut session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + let message = session + .append_turn_with_options(vec![ContentBlock::text("hi")], RunOptions::default()) + .await + .unwrap(); + + assert_eq!(message.text(), "hello"); + assert_eq!(session.metadata().turn_count, 1); + assert_eq!(session.metadata().status, SessionStatus::Idle); +} + +#[tokio::test] +async fn session_status_transitions_created_to_idle() { + let mock = MockRuntime::builder().text("done").build().unwrap(); + let mut session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + assert_eq!(session.metadata().status, SessionStatus::Created); + + let _message = session + .append_turn(vec![ContentBlock::text("go")]) + .await + .unwrap(); + + assert_eq!(session.metadata().status, SessionStatus::Idle); +} + +#[tokio::test] +async fn history_returns_committed_messages() { + let mock = MockRuntime::builder().text("response").build().unwrap(); + let mut session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + assert!(session.history().is_empty()); + + let _message = session + .append_turn(vec![ContentBlock::text("hello")]) + .await + .unwrap(); + + assert!( + !session.history().is_empty(), + "History should contain messages after a turn" + ); +} + +#[tokio::test] +async fn create_session_emits_session_started() { + let mock = MockRuntime::builder().text("hi").build().unwrap(); + + let session = mock + .runtime() + .create_session("test-session", mock.model()) + .unwrap(); + + // The SessionStarted event was emitted during creation. + // Verify session id follows the expected format. + assert!(session.id().as_str().starts_with("session-")); +} + +// ---- Task 4 permission tests ---- + +// -- PermissionDecision constructors -- + +#[test] +fn permission_decision_allow_constructor() { + let decision = PermissionDecision::allow(); + assert!(decision.allow); + assert!(decision.remember_as.is_none()); +} + +#[test] +fn permission_decision_deny_constructor() { + let decision = PermissionDecision::deny(); + assert!(!decision.allow); + assert!(decision.remember_as.is_none()); + assert!(decision.reason.is_none()); +} + +#[test] +fn permission_decision_allow_and_remember_constructor() { + let decision = PermissionDecision::allow_and_remember(PermissionRuleScope::Session); + assert!(decision.allow); + assert_eq!(decision.remember_as, Some(PermissionRuleScope::Session)); +} + +#[test] +fn permission_decision_deny_and_remember_constructor() { + let decision = PermissionDecision::deny_and_remember(PermissionRuleScope::Global); + assert!(!decision.allow); + assert_eq!(decision.remember_as, Some(PermissionRuleScope::Global)); +} + +#[test] +fn permission_decision_with_reason_keeps_the_rest_of_the_decision() { + let decision = PermissionDecision::deny_and_remember(PermissionRuleScope::Session) + .with_reason("no writes"); + assert!(!decision.allow); + assert_eq!(decision.remember_as, Some(PermissionRuleScope::Session)); + assert_eq!(decision.reason.as_deref(), Some("no writes")); +} + +// -- RuleStore -- + +#[test] +fn rule_store_empty_check_returns_none() { + let store = RuleStore::new(); + assert!(store.check("shell", None).is_none()); +} + +#[test] +fn rule_store_add_and_check_allow() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + assert_eq!(store.check("shell", None), Some(true)); +} + +#[test] +fn rule_store_add_and_check_deny() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow: false, + scope: PermissionRuleScope::Project, + reason: None, + }); + assert_eq!(store.check("shell", None), Some(false)); +} + +#[test] +fn rule_store_overwrite_replaces_rule() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + assert_eq!(store.check("shell", None), Some(true)); + + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow: false, + scope: PermissionRuleScope::Session, + reason: None, + }); + assert_eq!(store.check("shell", None), Some(false)); +} + +#[test] +fn rule_store_clear_scope_removes_matching_rules() { + let store = RuleStore::new(); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "read".to_owned(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Global, + reason: None, + }); + + store.clear_scope(PermissionRuleScope::Session); + + assert!(store.check("shell", None).is_none()); + assert_eq!(store.check("read", None), Some(true)); +} + +#[test] +fn rule_store_rules_returns_all_entries() { + let store = RuleStore::new(); + assert!(store.rules().is_empty()); + + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "shell".to_owned(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }); + store.add_rule(RememberedRule { + key: RuleKey { + tool_name: "read".to_owned(), + pattern: None, + }, + allow: false, + scope: PermissionRuleScope::Project, + reason: None, + }); + + assert_eq!(store.rules().len(), 2); +} + +// -- Session.resolve_permission -- + +#[tokio::test] +async fn resolve_permission_emits_event_and_sends_decision() { + let mock = MockRuntime::builder().text("hi").build().unwrap(); + let session = mock + .runtime() + .create_session("perm-test", mock.model()) + .unwrap(); + + let mut rx = session.subscribe(); + + // Simulate a pending permission by inserting directly. + let (tx, oneshot_rx) = tokio::sync::oneshot::channel(); + session.pending_permissions.insert( + "perm-1".to_owned(), + crate::session::permission::PendingPermissionEntry { + tool_call_id: "tc-1".to_owned(), + tool_name: "shell".to_owned(), + sender: tx, + }, + ); + + let decision = PermissionDecision::allow_and_remember(PermissionRuleScope::Session); + session.resolve_permission("perm-1", decision).unwrap(); + + // The oneshot should deliver the decision. + let received = oneshot_rx.await.unwrap(); + assert!(received.allow); + assert_eq!(received.remember_as, Some(PermissionRuleScope::Session)); + + // A PermissionResolved event should have been emitted. + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + let resolved = events.iter().find(|e| { + matches!( + e, + SessionEvent::PermissionResolved { + request_id, + outcome: PermissionOutcome::Allowed, + .. + } if request_id == "perm-1" + ) + }); + assert!( + resolved.is_some(), + "Expected PermissionResolved event, got: {events:?}" + ); + + // The rule should have been remembered. + let rules = session.remembered_rules(); + assert_eq!(rules.len(), 1); + assert!(rules[0].allow); +} + +/// Registers one pending permission on `session` and answers it with +/// `decision`, returning the rules the session remembered as a result. +async fn remembered_after( + session: &crate::session::Session, + decision: PermissionDecision, +) -> Vec { + let (tx, _rx) = tokio::sync::oneshot::channel(); + session.pending_permissions.insert( + "perm-1".to_owned(), + crate::session::permission::PendingPermissionEntry { + tool_call_id: "tc-1".to_owned(), + tool_name: "shell".to_owned(), + sender: tx, + }, + ); + session + .resolve_permission("perm-1", decision) + .expect("the pending permission should resolve"); + session.remembered_rules() +} + +#[tokio::test] +async fn a_remembered_refusal_keeps_the_reason_it_was_refused_with() { + // The approver answers once; the rule answers every time after that, so + // the reason has to survive into the rule or it is said only once. + let mock = MockRuntime::builder().text("hi").build().unwrap(); + let session = mock + .runtime() + .create_session("perm-reason", mock.model()) + .unwrap(); + + let rules = remembered_after( + &session, + PermissionDecision::deny_and_remember(PermissionRuleScope::Session) + .with_reason("this run does not allow writes"), + ) + .await; + + assert_eq!(rules.len(), 1); + assert_eq!( + rules[0].reason.as_deref(), + Some("this run does not allow writes") + ); +} + +#[tokio::test] +async fn a_remembered_allow_keeps_no_reason() { + let mock = MockRuntime::builder().text("hi").build().unwrap(); + let session = mock + .runtime() + .create_session("perm-reason-allow", mock.model()) + .unwrap(); + + let rules = remembered_after( + &session, + PermissionDecision::allow_and_remember(PermissionRuleScope::Session) + .with_reason("allowed for the session"), + ) + .await; + + assert_eq!(rules.len(), 1); + assert!(rules[0].allow); + assert_eq!( + rules[0].reason, None, + "an allowed call explains itself by happening" + ); +} + +#[tokio::test] +async fn resolve_permission_unknown_id_returns_error() { + let mock = MockRuntime::builder().text("hi").build().unwrap(); + let session = mock + .runtime() + .create_session("perm-test", mock.model()) + .unwrap(); + + let result = session.resolve_permission("nonexistent", PermissionDecision::deny()); + assert!(result.is_err()); +} + +#[derive(Clone)] +struct PromptingAuthorizer; + +#[async_trait] +impl crate::tool::ToolAuthorizer for PromptingAuthorizer { + async fn authorize( + &self, + _request: &crate::tool::ToolAuthorizationRequest, + ) -> Result { + Ok(crate::tool::ToolAuthorizationDecision::prompt( + "integration test prompt", + )) + } +} + +#[derive(Clone)] +struct InFlightProvider { + model: crate::ModelInfo, + turn: std::sync::Arc, +} + +#[async_trait] +impl crate::Provider for InFlightProvider { + fn descriptor(&self) -> crate::ProviderDescriptor { + crate::ProviderDescriptor::new(self.model.provider.clone()) + } + + async fn list_models(&self) -> Result, crate::ProviderError> { + Ok(vec![self.model.clone()]) + } + + async fn stream( + &self, + _request: crate::Request<'_>, + ) -> Result { + let turn = self.turn.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + let response = match turn { + 0 => crate::provider::Response { + id: unique_turn_id(), + model: self.model.id.clone(), + role: crate::Role::Assistant, + content: vec![crate::provider::ContentBlock::ToolUse { + id: "tool-1".to_string(), + name: "permission-test-tool".to_string(), + input: serde_json::json!({"input": "test"}), + }], + stop_reason: Some("tool_use".to_string()), + usage: None, + }, + _ => crate::provider::Response { + id: unique_turn_id(), + model: self.model.id.clone(), + role: crate::Role::Assistant, + content: vec![crate::provider::ContentBlock::text("final response")], + stop_reason: None, + usage: None, + }, + }; + Ok(crate::provider_event_stream_from_response(response)) + } +} + +#[derive(Clone)] +struct PromptTestTool; + +#[async_trait] +impl crate::tool::ToolDefinition for PromptTestTool { + fn descriptor(&self) -> crate::tool::ToolSpec { + crate::tool::ToolSpec::builder("permission-test-tool") + .description("Simple tool used for permission handle in-flight test") + .input_schema(serde_json::json!({ + "type": "object", + "properties": {} + })) + .build() + } +} + +#[async_trait] +impl crate::tool::ToolExecutor for PromptTestTool { + async fn execute( + &self, + _ctx: crate::tool::ParallelToolContext, + _input: serde_json::Value, + ) -> crate::tool::ToolResult { + Ok("tool-result".to_string()) + } +} + +fn unique_turn_id() -> String { + use std::fmt::Write as _; + + let mut out = String::new(); + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + let _ = write!(&mut out, "perm-flight-{nanos}"); + out +} + +#[tokio::test] +async fn resolve_permission_via_session_handle_while_append_turn_is_in_flight() { + use tokio::time::timeout; + + let model = crate::ModelInfo::new("mock-model", crate::BuiltinProvider::OpenAI); + let runtime = crate::Runtime::builder() + .with_tool_authorizer(PromptingAuthorizer) + .with_provider_instance(InFlightProvider { + model: model.clone(), + turn: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)), + }) + .with_policy(crate::RuntimePolicy::permissive()) + .build() + .unwrap(); + runtime.register_tool(PromptTestTool); + + let mut session = runtime + .create_session("permission-handle-flight", model.clone()) + .unwrap(); + let permission_handle = session.permission_handle(); + let mut events = session.subscribe(); + + let append = tokio::spawn(async move { + session + .append_turn(vec![ContentBlock::text("run permission test tool")]) + .await + }); + + let mut request_id = None; + for _ in 0..10 { + let event = timeout(std::time::Duration::from_millis(200), events.recv()) + .await + .expect("permission request should arrive") + .expect("session event stream should still be active"); + if let SessionEvent::PermissionRequested { + request_id: pending_id, + .. + } = event + { + request_id = Some(pending_id); + break; + } + } + + let request_id = request_id.expect("expected a PermissionRequested event"); + assert!(!append.is_finished()); + + permission_handle + .resolve_permission( + &request_id, + PermissionDecision::allow_and_remember(PermissionRuleScope::Session), + ) + .unwrap(); + + let result = append + .await + .expect("append turn task should complete") + .expect("append turn should succeed"); + assert_eq!(result.text(), "final response"); +} + +// ---- Task 8: Contract conformance integration tests ---- + +use async_trait::async_trait; +use serde_json::json; + +use crate::{ + provider::ProviderError, + test::MockToolCall, + tool::{ParallelToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSpec}, +}; + +struct EchoTool; + +#[async_trait] +impl ToolDefinition for EchoTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("echo_tool") + .description("Echo a canned result for testing") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .build() + } +} + +#[async_trait] +impl ToolExecutor for EchoTool { + async fn execute(&self, _ctx: ParallelToolContext, _input: serde_json::Value) -> ToolResult { + Ok("echoed".to_string()) + } +} + +#[tokio::test] +async fn full_session_lifecycle_produces_correct_event_stream() { + let mock = MockRuntime::builder() + .text("Hello, world!") + .build() + .unwrap(); + let mut session = mock + .runtime() + .create_session("lifecycle-test", mock.model()) + .unwrap(); + + let mut rx = session.subscribe(); + + let message = session + .append_turn(vec![ContentBlock::text("Hi there")]) + .await + .unwrap(); + + assert_eq!(message.text(), "Hello, world!"); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + // Verify UserMessage appears before AssistantMessageCompleted. + let user_pos = events + .iter() + .position(|e| matches!(e, SessionEvent::UserMessage { text } if text == "Hi there")); + let assistant_pos = events.iter().position(|e| { + matches!(e, SessionEvent::AssistantMessageCompleted { text } if text == "Hello, world!") + }); + + assert!( + user_pos.is_some(), + "Expected UserMessage event, got: {events:?}" + ); + assert!( + assistant_pos.is_some(), + "Expected AssistantMessageCompleted event, got: {events:?}" + ); + assert!( + user_pos.unwrap() < assistant_pos.unwrap(), + "UserMessage must precede AssistantMessageCompleted, positions: user={}, assistant={}", + user_pos.unwrap(), + assistant_pos.unwrap() + ); +} + +#[tokio::test] +async fn tool_call_session_produces_tool_lifecycle_events() { + let mock = MockRuntime::builder() + .tool_calls([MockToolCall::new("echo_tool", json!({}))]) + .text("tool work done") + .build() + .unwrap(); + mock.runtime().register_tool(EchoTool); + + let mut session = mock + .runtime() + .create_session("tool-test", mock.model()) + .unwrap(); + + let mut rx = session.subscribe(); + + let message = session + .append_turn(vec![ContentBlock::text("run the tool")]) + .await + .unwrap(); + + assert_eq!(message.text(), "tool work done"); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + let has_tool_started = events.iter().any( + |e| matches!(e, SessionEvent::ToolStarted { tool_name, .. } if tool_name == "echo_tool"), + ); + let has_tool_completed = events.iter().any(|e| { + matches!(e, SessionEvent::ToolCompleted { tool_call_id, .. } if tool_call_id == "tool-1") + }); + + assert!( + has_tool_started, + "Expected ToolStarted event for echo_tool, got: {events:?}" + ); + assert!( + has_tool_completed, + "Expected ToolCompleted event for tool-1, got: {events:?}" + ); +} + +#[derive(Clone)] +struct OverflowingToolProvider { + model: crate::ModelInfo, + turn: std::sync::Arc, +} + +#[async_trait] +impl crate::Provider for OverflowingToolProvider { + fn descriptor(&self) -> crate::ProviderDescriptor { + crate::ProviderDescriptor::new(self.model.provider.clone()) + } + + async fn list_models(&self) -> Result, crate::ProviderError> { + Ok(vec![self.model.clone()]) + } + + async fn stream( + &self, + _request: crate::Request<'_>, + ) -> Result { + let turn = self.turn.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + match turn { + 0 => Ok(buffered_provider_events(verbose_tool_turn_events( + &self.model.id, + 300, + ))), + _ => Ok(crate::provider_event_stream_from_response( + crate::provider::Response { + id: unique_turn_id(), + model: self.model.id.clone(), + role: crate::Role::Assistant, + content: vec![crate::provider::ContentBlock::text("tool run finished")], + stop_reason: None, + usage: None, + }, + )), + } + } +} + +fn buffered_provider_events( + events: Vec, +) -> crate::ProviderEventStream { + let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); + for event in events { + tx.send(Ok(event)) + .expect("session test provider receiver dropped unexpectedly"); + } + rx +} + +fn verbose_tool_turn_events( + model_id: &str, + delta_count: usize, +) -> Vec { + let mut events = vec![ + crate::provider::ProviderEvent::MessageStarted { + id: unique_turn_id(), + model: model_id.to_string(), + role: crate::Role::Assistant, + }, + crate::provider::ProviderEvent::ContentBlockStarted { + index: 0, + kind: crate::provider::ContentBlockStart::Text, + }, + ]; + + for index in 0..delta_count { + events.push(crate::provider::ProviderEvent::ContentBlockDelta { + index: 0, + delta: crate::provider::ContentBlockDelta::Text(format!("chunk-{index}")), + }); + } + + events.extend([ + crate::provider::ProviderEvent::ContentBlockStopped { index: 0 }, + crate::provider::ProviderEvent::ContentBlockStarted { + index: 1, + kind: crate::provider::ContentBlockStart::ToolUse { + id: "tool-1".to_string(), + name: "echo_tool".to_string(), + }, + }, + crate::provider::ProviderEvent::ContentBlockDelta { + index: 1, + delta: crate::provider::ContentBlockDelta::ToolUseInputJson("{}".to_string()), + }, + crate::provider::ProviderEvent::ContentBlockStopped { index: 1 }, + crate::provider::ProviderEvent::MessageDelta { + stop_reason: Some("tool_use".to_string()), + usage: None, + }, + crate::provider::ProviderEvent::MessageStopped, + ]); + + events +} + +#[tokio::test] +async fn session_preserves_tool_events_after_many_token_deltas() { + let model = crate::ModelInfo::new("mock-model", crate::BuiltinProvider::OpenAI); + let runtime = crate::Runtime::builder() + .with_provider_instance(OverflowingToolProvider { + model: model.clone(), + turn: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)), + }) + .with_policy(crate::RuntimePolicy::permissive()) + .build() + .unwrap(); + runtime.register_tool(EchoTool); + + let mut session = runtime + .create_session("overflow-tool-events", model.clone()) + .unwrap(); + let mut rx = session.subscribe(); + + let message = session + .append_turn(vec![ContentBlock::text("run the verbose tool turn")]) + .await + .unwrap(); + + assert_eq!(message.text(), "tool run finished"); + + let events: Vec<_> = std::iter::from_fn(|| rx.try_recv().ok()).collect(); + + let token_delta_count = events + .iter() + .filter(|event| matches!(event, SessionEvent::AssistantTokenDelta { .. })) + .count(); + let has_tool_started = events.iter().any( + |event| matches!(event, SessionEvent::ToolStarted { tool_name, .. } if tool_name == "echo_tool"), + ); + let has_tool_completed = events.iter().any(|event| { + matches!(event, SessionEvent::ToolCompleted { tool_call_id, .. } if tool_call_id == "tool-1") + }); + + assert!( + token_delta_count >= 300, + "Expected all token deltas to survive session mapping, got {token_delta_count} from {events:?}" + ); + assert!( + has_tool_started, + "Expected ToolStarted event after many token deltas, got: {events:?}" + ); + assert!( + has_tool_completed, + "Expected ToolCompleted event after many token deltas, got: {events:?}" + ); +} + +#[tokio::test] +async fn resume_session_restores_state() { + use crate::runtime::SqliteRuntimeStore; + use std::sync::atomic::{AtomicU64, Ordering}; + use std::time::{SystemTime, UNIX_EPOCH}; + + static NEXT_ID: AtomicU64 = AtomicU64::new(1); + let unique = NEXT_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + let store_path = + std::env::temp_dir().join(format!("mentra-session-resume-{timestamp}-{unique}.sqlite")); + let store = SqliteRuntimeStore::new(store_path); + let runtime_id = "resume-session-test"; + + let agent_id: String; + + // Phase 1: build a runtime, create a session, send a turn, then drop everything. + { + let mock = MockRuntime::builder() + .runtime_identifier(runtime_id) + .with_store(store.clone()) + .text("first response") + .build() + .unwrap(); + let mut session = mock + .runtime() + .create_session("resume-test", mock.model()) + .unwrap(); + + let _message = session + .append_turn(vec![ContentBlock::text("hello")]) + .await + .unwrap(); + + agent_id = session.agent_id().to_owned(); + + assert!( + !session.history().is_empty(), + "Session should have history after a turn" + ); + assert!( + !session.replay().items().is_empty(), + "Session transcript should be non-empty after a turn" + ); + // mock (and its Runtime) dropped here, releasing the agent lease. + } + + // Phase 2: build a fresh runtime with the same shared store, resume the agent. + let mock2 = MockRuntime::builder() + .runtime_identifier(runtime_id) + .with_store(store) + .build() + .unwrap(); + + let resumed_session = mock2.runtime().resume_session(&agent_id).unwrap(); + + assert!( + !resumed_session.replay().items().is_empty(), + "Resumed session should have a non-empty transcript" + ); +} + +#[tokio::test] +async fn failed_turn_emits_error_event() { + let mock = MockRuntime::builder() + .failure(ProviderError::InvalidResponse( + "provider exploded".to_string(), + )) + .build() + .unwrap(); + + let mut session = mock + .runtime() + .create_session("failure-test", mock.model()) + .unwrap(); + + let mut rx = session.subscribe(); + + let result = session + .append_turn(vec![ContentBlock::text("trigger failure")]) + .await; + + assert!(result.is_err(), "Expected append_turn to fail"); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + let has_error = events + .iter() + .any(|e| matches!(e, SessionEvent::Error { .. })); + assert!( + has_error, + "Expected Error event after failed turn, got: {events:?}" + ); +} + +#[tokio::test] +async fn all_session_events_from_turn_are_serializable_to_json() { + let mock = MockRuntime::builder() + .text("serializable response") + .build() + .unwrap(); + + let mut session = mock + .runtime() + .create_session("serde-test", mock.model()) + .unwrap(); + + let mut rx = session.subscribe(); + + let _message = session + .append_turn(vec![ContentBlock::text("check serde")]) + .await + .unwrap(); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + assert!( + !events.is_empty(), + "Expected at least one event from a turn" + ); + + for event in &events { + let json = serde_json::to_value(event) + .unwrap_or_else(|err| panic!("Failed to serialize event {event:?}: {err}")); + assert!( + json.get("type").is_some(), + "Serialized event missing 'type' tag: {json}" + ); + } +} + +// ---- Task 2A.2: File-edit event metadata ---- + +// The MockRuntime builder registers builtin tools (including `files`) automatically +// via Runtime::builder() -> RuntimeBuilder::new(true). + +use crate::agent::{AgentConfig, WorkspaceConfig}; + +/// Creates a unique temp directory and returns (base_dir, unique_suffix). +fn unique_test_base_dir(label: &str) -> std::path::PathBuf { + use std::time::{SystemTime, UNIX_EPOCH}; + let unique = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + let base_dir = std::env::temp_dir().join(format!("mentra-{label}-{unique}")); + std::fs::create_dir_all(&base_dir).unwrap(); + std::fs::canonicalize(&base_dir).unwrap_or(base_dir) +} + +#[tokio::test] +async fn files_tool_create_emits_tool_progress_with_file_op_metadata() { + let base_dir = unique_test_base_dir("file-op-test"); + let target_path = base_dir.join("hello.txt"); + + let mock = MockRuntime::builder() + .tool_calls([MockToolCall::new( + "files", + json!({ + "operations": [ + { + "op": "create", + "path": target_path.to_str().unwrap(), + "content": "hello world\n" + } + ] + }), + )]) + .text("file created") + .build() + .unwrap(); + + // Use create_session_with_config so the agent's base_dir covers the temp dir, + // satisfying the runtime policy write-root check. + let agent_config = AgentConfig { + workspace: WorkspaceConfig { + base_dir: base_dir.clone(), + auto_route_shell: false, + }, + ..AgentConfig::default() + }; + let mut session = mock + .runtime() + .create_session_with_config("file-op-test", mock.model(), agent_config) + .unwrap(); + + let mut rx = session.subscribe(); + + let message = session + .append_turn(vec![ContentBlock::text("create a file")]) + .await + .unwrap(); + + assert_eq!(message.text(), "file created"); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + // Verify that at least one ToolProgress event was emitted with a + // "file_op:" prefix indicating the create operation. + let file_op_progress = events.iter().find(|e| { + matches!( + e, + SessionEvent::ToolProgress { progress, .. } + if progress.starts_with("file_op: create ") + ) + }); + + assert!( + file_op_progress.is_some(), + "Expected ToolProgress event with 'file_op: create ...' metadata, got: {events:?}" + ); + + // Confirm the created file actually exists on disk. + assert!( + target_path.exists(), + "Expected file to exist at {target_path:?}" + ); + + // Clean up. + let _ = std::fs::remove_dir_all(&base_dir); +} + +#[tokio::test] +async fn files_tool_set_emits_tool_progress_with_file_op_metadata() { + let base_dir = unique_test_base_dir("file-set-test"); + let target_path = base_dir.join("target.txt"); + std::fs::write(&target_path, "original content\n").unwrap(); + + let mock = MockRuntime::builder() + .tool_calls([MockToolCall::new( + "files", + json!({ + "operations": [ + { + "op": "set", + "path": target_path.to_str().unwrap(), + "content": "updated content\n" + } + ] + }), + )]) + .text("file updated") + .build() + .unwrap(); + + let agent_config = AgentConfig { + workspace: WorkspaceConfig { + base_dir: base_dir.clone(), + auto_route_shell: false, + }, + ..AgentConfig::default() + }; + let mut session = mock + .runtime() + .create_session_with_config("file-set-test", mock.model(), agent_config) + .unwrap(); + + let mut rx = session.subscribe(); + + let message = session + .append_turn(vec![ContentBlock::text("update the file")]) + .await + .unwrap(); + + assert_eq!(message.text(), "file updated"); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + let file_op_progress = events.iter().find(|e| { + matches!( + e, + SessionEvent::ToolProgress { progress, .. } + if progress.starts_with("file_op: set ") + ) + }); + + assert!( + file_op_progress.is_some(), + "Expected ToolProgress event with 'file_op: set ...' metadata, got: {events:?}" + ); + + let _ = std::fs::remove_dir_all(&base_dir); +} + +#[tokio::test] +async fn files_tool_read_does_not_emit_file_op_progress() { + let base_dir = unique_test_base_dir("file-read-test"); + let target_path = base_dir.join("read_me.txt"); + std::fs::write(&target_path, "some content\n").unwrap(); + + let mock = MockRuntime::builder() + .tool_calls([MockToolCall::new( + "files", + json!({ + "operations": [ + { + "op": "read", + "path": target_path.to_str().unwrap() + } + ] + }), + )]) + .text("read done") + .build() + .unwrap(); + + let agent_config = AgentConfig { + workspace: WorkspaceConfig { + base_dir: base_dir.clone(), + auto_route_shell: false, + }, + ..AgentConfig::default() + }; + let mut session = mock + .runtime() + .create_session_with_config("file-read-test", mock.model(), agent_config) + .unwrap(); + + let mut rx = session.subscribe(); + + let message = session + .append_turn(vec![ContentBlock::text("read the file")]) + .await + .unwrap(); + + assert_eq!(message.text(), "read done"); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + // Read operations must NOT emit file_op: progress events. + let file_op_progress = events.iter().find(|e| { + matches!( + e, + SessionEvent::ToolProgress { progress, .. } + if progress.starts_with("file_op:") + ) + }); + + assert!( + file_op_progress.is_none(), + "Read operation should not emit file_op progress, but got: {file_op_progress:?}" + ); + + let _ = std::fs::remove_dir_all(&base_dir); +} + +// ---- Task 2A.4: Compaction continuity ---- + +#[tokio::test] +async fn compaction_events_appear_in_session_stream_and_session_continues() { + // Use a very low auto_compact_threshold so compaction triggers after the first turn. + // The mock provider needs: turn 1 response, compaction summary, turn 2 response. + let transcript_dir = unique_test_base_dir("compact-session"); + let agent_config = AgentConfig { + compaction: crate::agent::CompactionConfig { + auto_compact_threshold_tokens: Some(1), + transcript_dir: transcript_dir.clone(), + ..Default::default() + }, + ..AgentConfig::default() + }; + + let mock = MockRuntime::builder() + .text("first response") + .text("compaction summary") // consumed by the compaction summarizer + .text("second response") + .build() + .unwrap(); + + let mut session = mock + .runtime() + .create_session_with_config("compact-test", mock.model(), agent_config) + .unwrap(); + + let mut rx = session.subscribe(); + + // Turn 1: triggers compaction because threshold is 1 token. + let msg1 = session + .append_turn(vec![ContentBlock::text("first turn")]) + .await + .unwrap(); + assert_eq!(msg1.text(), "first response"); + + // Turn 2: session continues coherently after compaction. + let msg2 = session + .append_turn(vec![ContentBlock::text("second turn")]) + .await + .unwrap(); + assert_eq!(msg2.text(), "second response"); + assert_eq!(session.metadata().turn_count, 2); + assert_eq!(session.metadata().status, SessionStatus::Idle); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + let has_compaction_started = events + .iter() + .any(|e| matches!(e, SessionEvent::CompactionStarted { .. })); + let has_compaction_completed = events + .iter() + .any(|e| matches!(e, SessionEvent::CompactionCompleted { .. })); + + assert!( + has_compaction_started, + "Expected CompactionStarted event, got: {events:?}" + ); + assert!( + has_compaction_completed, + "Expected CompactionCompleted event, got: {events:?}" + ); + + // Verify ordering: CompactionStarted before CompactionCompleted. + let started_pos = events + .iter() + .position(|e| matches!(e, SessionEvent::CompactionStarted { .. })) + .unwrap(); + let completed_pos = events + .iter() + .position(|e| matches!(e, SessionEvent::CompactionCompleted { .. })) + .unwrap(); + assert!( + started_pos < completed_pos, + "CompactionStarted (pos {started_pos}) must precede CompactionCompleted (pos {completed_pos})" + ); + + // Verify the second turn's assistant response appears after compaction. + // Note: The UserMessage for the second turn is emitted before agent.send(), + // so it may precede compaction events (compaction triggers during send). + // But the AssistantMessageCompleted for turn 2 must come after compaction. + let second_assistant = events.iter().position(|e| { + matches!(e, SessionEvent::AssistantMessageCompleted { text } if text == "second response") + }); + assert!( + second_assistant.is_some(), + "Expected second turn AssistantMessageCompleted after compaction" + ); + assert!( + second_assistant.unwrap() > completed_pos, + "Second turn assistant response must appear after compaction completed" + ); + + let _ = std::fs::remove_dir_all(&transcript_dir); +} + +// ---- Task 2A.6: Session resume continuity ---- + +#[tokio::test] +async fn resume_session_with_permission_rules_restores_rules() { + use crate::runtime::{PermissionRuleStore, SqliteRuntimeStore}; + use std::sync::Arc; + + let unique = unique_test_base_dir("resume-rules"); + let store_path = unique.join("runtime.sqlite"); + let store = SqliteRuntimeStore::new(&store_path); + let runtime_id = "resume-rules-test"; + + let session_id_str: String; + let agent_id: String; + + // Phase 1: Create session, add a permission rule, persist it. + { + let mock = MockRuntime::builder() + .runtime_identifier(runtime_id) + .with_store(store.clone()) + .text("hello") + .build() + .unwrap(); + + let mut session = mock + .runtime() + .create_session("resume-rules", mock.model()) + .unwrap(); + + session.set_permission_store(Arc::new(store.clone()) as Arc); + session_id_str = session.id().as_str().to_owned(); + agent_id = session.agent_id().to_owned(); + + // Simulate a permission decision that gets remembered. + let permission_handle = session.permission_handle(); + let (tx, _rx) = tokio::sync::oneshot::channel(); + session.pending_permissions.insert( + "perm-r1".to_owned(), + crate::session::permission::PendingPermissionEntry { + tool_call_id: "tc-r1".to_owned(), + tool_name: "shell".to_owned(), + sender: tx, + }, + ); + permission_handle + .resolve_permission( + "perm-r1", + PermissionDecision::allow_and_remember(PermissionRuleScope::Session), + ) + .unwrap(); + + // Verify rule is in memory. + assert_eq!(session.remembered_rules().len(), 1); + + // Submit a turn so the session has history. + let _msg = session + .append_turn(vec![ContentBlock::text("hi")]) + .await + .unwrap(); + + // Session + runtime dropped here. + } + + // Phase 2: Resume session, attach same store, load persisted rules. + let mock2 = MockRuntime::builder() + .runtime_identifier(runtime_id) + .with_store(store.clone()) + .text("resumed response") + .build() + .unwrap(); + + let mut resumed = mock2.runtime().resume_session(&agent_id).unwrap(); + resumed.set_permission_store(Arc::new(store.clone()) as Arc); + + // Load the persisted rules using the original session id. + // Note: resume_session creates a new SessionId, so we must load from the + // original session id that was used when persisting. This tests the store + // directly. + let loaded_rules = store.load_rules(&session_id_str, None).unwrap(); + assert_eq!( + loaded_rules.len(), + 1, + "Expected 1 persisted rule, got: {loaded_rules:?}" + ); + assert!(loaded_rules[0].allow); + assert_eq!(loaded_rules[0].key.tool_name, "shell"); + assert!( + store.load_rules("perm-r1", None).unwrap().is_empty(), + "Expected rules to be saved under session id, not permission request id" + ); + + // Verify resumed session has intact transcript. + assert!( + !resumed.replay().items().is_empty(), + "Resumed session should have non-empty transcript" + ); + + // Verify resumed session can accept new turns. + let msg = resumed + .append_turn(vec![ContentBlock::text("after resume")]) + .await + .unwrap(); + assert_eq!(msg.text(), "resumed response"); + assert_eq!(resumed.metadata().turn_count, 1); + + let _ = std::fs::remove_dir_all(&unique); +} + +#[tokio::test] +async fn session_set_model_updates_active_and_persisted_model() { + let unique = unique_test_base_dir("set-model"); + let store_path = unique.join("runtime.sqlite"); + let store = SqliteRuntimeStore::new(&store_path); + let mock = MockRuntime::builder() + .with_store(store.clone()) + .text("hello") + .build() + .unwrap(); + + let mut session = mock + .runtime() + .create_session("switch-model", mock.model()) + .unwrap(); + let agent_id = session.agent_id().to_owned(); + let updated_model = crate::ModelInfo::new("switched-model", crate::BuiltinProvider::OpenAI); + + session.set_model(updated_model.clone()).unwrap(); + + assert_eq!(session.metadata().model, "switched-model"); + + let loaded = store + .load_agent(&agent_id) + .unwrap() + .expect("persisted agent should exist"); + assert_eq!(loaded.record.model, "switched-model"); + assert_eq!(loaded.record.provider_id, updated_model.provider); + + let _ = std::fs::remove_dir_all(&unique); +} + +// ---- Task 2A.7: Error handling and recovery ---- + +#[tokio::test] +async fn error_recovery_session_accepts_turn_after_failure() { + // Script: first turn fails, second turn succeeds. + let mock = MockRuntime::builder() + .failure(ProviderError::InvalidResponse( + "transient glitch".to_string(), + )) + .text("recovered successfully") + .build() + .unwrap(); + + let mut session = mock + .runtime() + .create_session("error-recovery", mock.model()) + .unwrap(); + + let mut rx = session.subscribe(); + + // Turn 1: fails. + let result = session + .append_turn(vec![ContentBlock::text("will fail")]) + .await; + assert!(result.is_err()); + assert!(matches!( + session.metadata().status, + SessionStatus::Failed(_) + )); + + // Turn 2: succeeds, proving session is recoverable. + let msg = session + .append_turn(vec![ContentBlock::text("retry")]) + .await + .unwrap(); + assert_eq!(msg.text(), "recovered successfully"); + assert_eq!(session.metadata().status, SessionStatus::Idle); + assert_eq!(session.metadata().turn_count, 1); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + // Verify error event from first turn. + let error_event = events.iter().find(|e| { + matches!( + e, + SessionEvent::Error { message, .. } + if message.contains("transient glitch") + ) + }); + assert!( + error_event.is_some(), + "Expected Error event containing 'transient glitch', got: {events:?}" + ); + + // Verify successful second turn events follow the error. + let error_pos = events + .iter() + .position(|e| matches!(e, SessionEvent::Error { .. })) + .unwrap(); + let second_assistant = events.iter().position(|e| { + matches!(e, SessionEvent::AssistantMessageCompleted { text } if text == "recovered successfully") + }); + assert!( + second_assistant.is_some() && second_assistant.unwrap() > error_pos, + "Second turn assistant message must appear after error event" + ); +} + +#[tokio::test] +async fn tool_execution_error_emits_tool_completed_with_is_error() { + use crate::tool::{ParallelToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSpec}; + + struct FailingTool; + + #[async_trait] + impl ToolDefinition for FailingTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("failing_tool") + .description("Always fails") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .build() + } + } + + #[async_trait] + impl ToolExecutor for FailingTool { + async fn execute( + &self, + _ctx: ParallelToolContext, + _input: serde_json::Value, + ) -> ToolResult { + Err("tool execution failed".to_string()) + } + } + + let mock = MockRuntime::builder() + .tool_calls([MockToolCall::new("failing_tool", json!({}))]) + .text("continued after tool failure") + .build() + .unwrap(); + mock.runtime().register_tool(FailingTool); + + let mut session = mock + .runtime() + .create_session("tool-fail-test", mock.model()) + .unwrap(); + + let mut rx = session.subscribe(); + + let msg = session + .append_turn(vec![ContentBlock::text("run failing tool")]) + .await + .unwrap(); + + // The session should continue even though the tool failed. + assert_eq!(msg.text(), "continued after tool failure"); + assert_eq!(session.metadata().status, SessionStatus::Idle); + + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + // Verify a ToolCompleted event with is_error = true appears. + let tool_error = events + .iter() + .find(|e| matches!(e, SessionEvent::ToolCompleted { is_error: true, .. })); + assert!( + tool_error.is_some(), + "Expected ToolCompleted with is_error=true, got: {events:?}" + ); +} + +// ---- Task 2A.8: Full scenario integration test ---- + +#[tokio::test] +async fn full_scenario_prompt_shell_file_events_end_to_end() { + let base_dir = unique_test_base_dir("scenario-e2e"); + let target_file = base_dir.join("scenario_output.txt"); + + // Script the full scenario: + // 1. Text response to initial prompt + // 2. Tool call (echo_tool simulating shell) + File create operation + // 3. Final text response + let mock = MockRuntime::builder() + .text("I will help you with that.") + .tool_calls([MockToolCall::new("echo_tool", json!({}))]) + .tool_calls([MockToolCall::new( + "files", + json!({ + "operations": [ + { + "op": "create", + "path": target_file.to_str().unwrap(), + "content": "scenario test output\n" + } + ] + }), + )]) + .text("All tasks completed successfully.") + .build() + .unwrap(); + mock.runtime().register_tool(EchoTool); + + let agent_config = AgentConfig { + workspace: WorkspaceConfig { + base_dir: base_dir.clone(), + auto_route_shell: false, + }, + ..AgentConfig::default() + }; + let mut session = mock + .runtime() + .create_session_with_config("scenario-test", mock.model(), agent_config) + .unwrap(); + + let mut rx = session.subscribe(); + + // Turn 1: Simple text response. + let msg1 = session + .append_turn(vec![ContentBlock::text("Help me set up a project")]) + .await + .unwrap(); + assert_eq!(msg1.text(), "I will help you with that."); + assert_eq!(session.metadata().turn_count, 1); + + // Turn 2: Tool calls (echo_tool + file create). + let msg2 = session + .append_turn(vec![ContentBlock::text("Create a file and run a command")]) + .await + .unwrap(); + assert_eq!(msg2.text(), "All tasks completed successfully."); + assert_eq!(session.metadata().turn_count, 2); + assert_eq!(session.metadata().status, SessionStatus::Idle); + + // Collect all events. + let mut events = Vec::new(); + while let Ok(event) = rx.try_recv() { + events.push(event); + } + + // Verify event ordering across the full scenario. + let event_types: Vec<&str> = events + .iter() + .map(|e| match e { + SessionEvent::UserMessage { .. } => "user_message", + SessionEvent::AssistantTokenDelta { .. } => "token_delta", + SessionEvent::AssistantReasoningDelta { .. } => "reasoning_delta", + SessionEvent::AssistantMessageCompleted { .. } => "assistant_completed", + SessionEvent::ToolQueued { .. } => "tool_queued", + SessionEvent::ToolStarted { .. } => "tool_started", + SessionEvent::ToolProgress { .. } => "tool_progress", + SessionEvent::ToolCompleted { .. } => "tool_completed", + SessionEvent::PermissionRequested { .. } => "perm_requested", + SessionEvent::PermissionResolved { .. } => "perm_resolved", + SessionEvent::TaskUpdated { .. } => "task_updated", + SessionEvent::CompactionStarted { .. } => "compaction_started", + SessionEvent::CompactionCompleted { .. } => "compaction_completed", + SessionEvent::MemoryUpdated { .. } => "memory_updated", + SessionEvent::Branched { .. } => "branched", + SessionEvent::Notice { .. } => "notice", + SessionEvent::RetryAttempt { .. } => "retry_attempt", + SessionEvent::Error { .. } => "error", + SessionEvent::UsageReport { .. } => "usage_report", + SessionEvent::SessionStarted { .. } => "session_started", + }) + .collect(); + + // Verify we got all expected event categories. + assert!( + event_types.contains(&"user_message"), + "Missing user_message events: {event_types:?}" + ); + assert!( + event_types.contains(&"assistant_completed"), + "Missing assistant_completed events: {event_types:?}" + ); + assert!( + event_types.contains(&"tool_started"), + "Missing tool_started events: {event_types:?}" + ); + assert!( + event_types.contains(&"tool_completed"), + "Missing tool_completed events: {event_types:?}" + ); + + // Verify there are exactly 2 UserMessage events (one per turn). + let user_msg_count = events + .iter() + .filter(|e| matches!(e, SessionEvent::UserMessage { .. })) + .count(); + assert_eq!(user_msg_count, 2, "Expected 2 UserMessage events"); + + // Verify there are exactly 2 AssistantMessageCompleted events. + let assistant_count = events + .iter() + .filter(|e| matches!(e, SessionEvent::AssistantMessageCompleted { .. })) + .count(); + assert_eq!( + assistant_count, 2, + "Expected 2 AssistantMessageCompleted events" + ); + + // Verify the file was actually created on disk. + assert!( + target_file.exists(), + "Expected scenario output file at {target_file:?}" + ); + + // Verify a file_op progress event was emitted for the create operation. + let has_file_op = events.iter().any(|e| { + matches!( + e, + SessionEvent::ToolProgress { progress, .. } + if progress.starts_with("file_op: create ") + ) + }); + assert!( + has_file_op, + "Expected file_op progress event for create, events: {event_types:?}" + ); + + // Verify echo_tool execution events appear. + let has_echo_tool = events.iter().any(|e| { + matches!( + e, + SessionEvent::ToolStarted { tool_name, .. } + if tool_name == "echo_tool" + ) + }); + assert!( + has_echo_tool, + "Expected ToolStarted for echo_tool, events: {event_types:?}" + ); + + let _ = std::fs::remove_dir_all(&base_dir); +} + +// ---- Task 1: project_id support in PermissionRuleStore ---- + +#[tokio::test] +async fn load_rules_with_project_id_returns_all_applicable_scopes() { + use crate::runtime::{PermissionRuleStore, SqliteRuntimeStore}; + use crate::session::event::PermissionRuleScope; + use crate::session::permission::{RememberedRule, RuleKey}; + + // Create an isolated temp store. + let store_path = std::env::temp_dir().join(format!( + "mentra-session-perm-scopes-{}.sqlite", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() + )); + let store = SqliteRuntimeStore::new(&store_path); + + let project_id = "test-project-42"; + let session_id = "session-main"; + let other_session_id = "session-other"; + + // Session-scoped rule: belongs to session-main only. + let session_rule = RememberedRule { + key: RuleKey { + tool_name: "shell".to_string(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }; + + // Project-scoped rule: saved under a different session but same project_id. + let project_rule = RememberedRule { + key: RuleKey { + tool_name: "file_write".to_string(), + pattern: Some("/workspace/*".to_string()), + }, + allow: true, + scope: PermissionRuleScope::Project, + reason: None, + }; + + // Global-scoped rule: saved under yet another session, no project. + let global_rule = RememberedRule { + key: RuleKey { + tool_name: "network".to_string(), + pattern: None, + }, + allow: false, + scope: PermissionRuleScope::Global, + reason: None, + }; + + // Save session-scoped rule under session-main with project_id. + store + .save_rules(session_id, Some(project_id), &[session_rule]) + .expect("save session rule"); + + // Save project-scoped rule under a different session but same project_id. + store + .save_rules(other_session_id, Some(project_id), &[project_rule]) + .expect("save project rule"); + + // Save global-scoped rule under another unrelated session (no project). + store + .save_rules("session-global-only", None, &[global_rule]) + .expect("save global rule"); + + // Load for session-main with project_id: should return all three scopes. + let loaded = store + .load_rules(session_id, Some(project_id)) + .expect("load rules"); + + assert_eq!( + loaded.len(), + 3, + "Expected session + project + global rules (3 total), got: {loaded:?}" + ); + + let has_session = loaded + .iter() + .any(|r| r.key.tool_name == "shell" && r.scope == PermissionRuleScope::Session); + let has_project = loaded + .iter() + .any(|r| r.key.tool_name == "file_write" && r.scope == PermissionRuleScope::Project); + let has_global = loaded + .iter() + .any(|r| r.key.tool_name == "network" && r.scope == PermissionRuleScope::Global); + + assert!(has_session, "Session-scoped rule should be present"); + assert!(has_project, "Project-scoped rule should be present"); + assert!(has_global, "Global-scoped rule should be present"); + + // Loading without project_id should only return session + global scopes. + let loaded_no_project = store + .load_rules(session_id, None) + .expect("load rules without project_id"); + + assert_eq!( + loaded_no_project.len(), + 2, + "Without project_id, expect only session + global rules, got: {loaded_no_project:?}" + ); + assert!( + loaded_no_project + .iter() + .any(|r| r.scope == PermissionRuleScope::Session), + "Session rule should still be present" + ); + assert!( + loaded_no_project + .iter() + .any(|r| r.scope == PermissionRuleScope::Global), + "Global rule should still be present" + ); + assert!( + !loaded_no_project + .iter() + .any(|r| r.scope == PermissionRuleScope::Project), + "Project rule should NOT be present when no project_id given" + ); + + let _ = std::fs::remove_file(&store_path); +} + +// ---- Task 3: Cross-session permission inheritance integration tests ---- + +#[tokio::test] +async fn project_scoped_rules_are_visible_across_sessions() { + use crate::runtime::{PermissionRuleStore, SqliteRuntimeStore}; + use crate::session::event::PermissionRuleScope; + use crate::session::permission::{RememberedRule, RuleKey}; + + let store_path = std::env::temp_dir().join(format!( + "mentra-cross-session-project-{}.sqlite", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() + )); + let store = SqliteRuntimeStore::new(&store_path); + + let project_id = "my-project"; + let session_1 = "session-1"; + let session_2 = "session-2"; + + let project_rule = RememberedRule { + key: RuleKey { + tool_name: "file_write".to_string(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Project, + reason: None, + }; + + // Session-1 saves a project-scoped rule for "my-project". + store + .save_rules(session_1, Some(project_id), &[project_rule]) + .expect("save project-scoped rule via session-1"); + + // Session-2 loads rules for the same project — project-scoped rule must be visible. + let loaded = store + .load_rules(session_2, Some(project_id)) + .expect("load rules for session-2"); + + let has_project_rule = loaded + .iter() + .any(|r| r.key.tool_name == "file_write" && r.scope == PermissionRuleScope::Project); + + assert!( + has_project_rule, + "Project-scoped rule saved by session-1 should be visible to session-2 under the same project_id, got: {loaded:?}" + ); + + let _ = std::fs::remove_file(&store_path); +} + +#[tokio::test] +async fn global_scoped_rules_are_visible_to_all_sessions() { + use crate::runtime::{PermissionRuleStore, SqliteRuntimeStore}; + use crate::session::event::PermissionRuleScope; + use crate::session::permission::{RememberedRule, RuleKey}; + + let store_path = std::env::temp_dir().join(format!( + "mentra-cross-session-global-{}.sqlite", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() + )); + let store = SqliteRuntimeStore::new(&store_path); + + let session_1 = "session-1"; + let session_2 = "session-2"; + + let global_rule = RememberedRule { + key: RuleKey { + tool_name: "network".to_string(), + pattern: None, + }, + allow: false, + scope: PermissionRuleScope::Global, + reason: None, + }; + + // Session-1 saves a global rule (no project_id). + store + .save_rules(session_1, None, &[global_rule]) + .expect("save global-scoped rule via session-1"); + + // Session-2 loads rules for a completely different project — global rule must be visible. + let loaded = store + .load_rules(session_2, Some("other-project")) + .expect("load rules for session-2 with other-project"); + + let has_global_rule = loaded + .iter() + .any(|r| r.key.tool_name == "network" && r.scope == PermissionRuleScope::Global); + + assert!( + has_global_rule, + "Global-scoped rule saved by session-1 should be visible to session-2 regardless of project, got: {loaded:?}" + ); + + let _ = std::fs::remove_file(&store_path); +} + +#[tokio::test] +async fn session_scoped_rules_are_not_visible_to_other_sessions() { + use crate::runtime::{PermissionRuleStore, SqliteRuntimeStore}; + use crate::session::event::PermissionRuleScope; + use crate::session::permission::{RememberedRule, RuleKey}; + + let store_path = std::env::temp_dir().join(format!( + "mentra-cross-session-isolation-{}.sqlite", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() + )); + let store = SqliteRuntimeStore::new(&store_path); + + let project_id = "shared-project"; + let session_1 = "session-1"; + let session_2 = "session-2"; + + let session_rule = RememberedRule { + key: RuleKey { + tool_name: "shell".to_string(), + pattern: None, + }, + allow: true, + scope: PermissionRuleScope::Session, + reason: None, + }; + + // Session-1 saves a session-scoped rule. + store + .save_rules(session_1, Some(project_id), &[session_rule]) + .expect("save session-scoped rule via session-1"); + + // Session-2 loads rules for the same project — session-1's session-scoped rule must NOT appear. + let loaded = store + .load_rules(session_2, Some(project_id)) + .expect("load rules for session-2"); + + let has_session_1_rule = loaded + .iter() + .any(|r| r.key.tool_name == "shell" && r.scope == PermissionRuleScope::Session); + + assert!( + !has_session_1_rule, + "Session-scoped rule from session-1 must NOT be visible to session-2 (session isolation), got: {loaded:?}" + ); + + let _ = std::fs::remove_file(&store_path); +} + +// ---- Host-facing subagent spawning ---- + +/// Waits for the terminal `TaskUpdated` a detached subagent broadcasts when it +/// finishes, skipping the `Spawned` notice. +async fn next_subagent_outcome( + events: &mut crate::session::SessionEventReceiver, +) -> (TaskLifecycleStatus, Option) { + let deadline = std::time::Duration::from_secs(10); + tokio::time::timeout(deadline, async { + loop { + match events.recv().await.unwrap() { + SessionEvent::TaskUpdated { + kind: TaskKind::Subagent, + status, + detail, + .. + } if status != TaskLifecycleStatus::Spawned => return (status, detail), + _ => continue, + } + } + }) + .await + .expect("a detached subagent must broadcast a terminal status") +} + +#[tokio::test] +async fn spawn_subagent_runs_on_default_options() { + // Host-facing and deliberately uninherited: this spawn is not made from + // inside a parent run, so there are no in-flight bounds to share. The + // opposite of the model-facing `task` intrinsic, which can only run inside + // one and always inherits. + let mock = MockRuntime::builder() + .text("subagent answer") + .build() + .unwrap(); + let mut session = mock + .runtime() + .create_session("host-spawn", mock.model()) + .unwrap(); + let mut events = session.subscribe(); + + session.spawn_subagent("research", "go").await.unwrap(); + + let (status, detail) = next_subagent_outcome(&mut events).await; + assert_eq!(status, TaskLifecycleStatus::Finished); + assert_eq!(detail.as_deref(), Some("subagent answer")); +} + +#[tokio::test] +async fn spawn_subagent_with_options_puts_the_subagent_under_them() { + // The opt-in variant is how a host reaches `RunOptions::child` for a + // detached subagent. Tripped before the spawn so the outcome does not depend + // on winning a race with the provider. + let mock = MockRuntime::builder() + .text("never reached") + .build() + .unwrap(); + let mut session = mock + .runtime() + .create_session("host-spawn-bounded", mock.model()) + .unwrap(); + let mut events = session.subscribe(); + + let cancellation = CancellationToken::default(); + let turn_options = RunOptions { + cancellation: Some(cancellation.clone()), + ..RunOptions::default() + }; + cancellation.cancel(); + + session + .spawn_subagent_with_options("research", "go", turn_options.child()) + .await + .unwrap(); + + let (status, _) = next_subagent_outcome(&mut events).await; + assert_eq!( + status, + TaskLifecycleStatus::Failed, + "the supplied options must reach the detached run" + ); + assert_eq!( + mock.recorded_requests().await.len(), + 0, + "the cancelled subagent never reached the provider" + ); +} diff --git a/vendor/mentra/src/session/tests/terminal_output.rs b/vendor/mentra/src/session/tests/terminal_output.rs new file mode 100644 index 0000000..1a452c9 --- /dev/null +++ b/vendor/mentra/src/session/tests/terminal_output.rs @@ -0,0 +1,568 @@ +//! Tests for [`Session::append_turn_to_output`]: the typed value it hands +//! back, what the session stream says around a typed turn, and what a typed +//! turn that fails leaves behind. + +use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, +}; + +use async_trait::async_trait; +use serde::Deserialize; +use serde_json::{Value, json}; + +use crate::{ + BuiltinProvider, ContentBlock, ModelInfo, Provider, ProviderDescriptor, ProviderError, + ProviderEventStream, Request, Role, Runtime, TerminalOutputSpec, ToolChoice, + provider::Response, + provider_event_stream_from_response, + runtime::RunOptions, + session::{Session, SessionEvent, SessionStatus}, + tool::{ + ToolContext, ToolDefinition, ToolDurability, ToolExecutor, ToolResult, ToolSideEffectLevel, + ToolSpec, + }, +}; + +/// What the model writes when it declines the forced tool and answers in prose. +const PLAIN_ANSWER: &str = "plain answer"; + +#[derive(Debug, Deserialize, PartialEq, Eq)] +struct Report { + answer: u64, + evidence: Vec, +} + +/// Plays the model for a terminal-output run: when the request forces one +/// tool, it calls exactly that tool, so a test never has to know the tool name +/// `run_to_output` generates for the run. Without a forced choice — an +/// ordinary turn on the same session — it answers in prose. +#[derive(Clone)] +struct ForcedToolProvider { + model: ModelInfo, + /// Prose the model writes alongside the terminal call. + preface: Option, + /// Input the model sends to the forced tool. `None` makes it ignore the + /// forced choice and answer in prose instead. + payload: Option, + calls: Arc, +} + +impl ForcedToolProvider { + fn new(payload: Option) -> Self { + Self { + model: ModelInfo::new("typed-output-model", BuiltinProvider::Anthropic), + preface: None, + payload, + calls: Arc::new(AtomicUsize::new(0)), + } + } + + fn answering(payload: Value) -> Self { + Self::new(Some(payload)) + } + + fn ignoring_the_forced_tool() -> Self { + Self::new(None) + } + + fn with_preface(mut self, preface: &str) -> Self { + self.preface = Some(preface.to_string()); + self + } +} + +#[async_trait] +impl Provider for ForcedToolProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.model.provider.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(vec![self.model.clone()]) + } + + async fn stream(&self, request: Request<'_>) -> Result { + let call = self.calls.fetch_add(1, Ordering::SeqCst); + let forced = match request.tool_choice.clone() { + Some(ToolChoice::Tool { name }) => Some(name), + _ => None, + }; + + let (content, stop_reason) = match (forced, self.payload.clone()) { + (Some(name), Some(payload)) => { + let mut blocks = Vec::new(); + if let Some(preface) = &self.preface { + blocks.push(ContentBlock::text(preface.clone())); + } + blocks.push(ContentBlock::ToolUse { + id: format!("terminal-call-{call}"), + name, + input: payload, + }); + (blocks, Some("tool_use".to_string())) + } + _ => (vec![ContentBlock::text(PLAIN_ANSWER)], None), + }; + + Ok(provider_event_stream_from_response(Response { + id: format!("message-{call}-{}", unique_suffix()), + model: self.model.id.clone(), + role: Role::Assistant, + content, + stop_reason, + usage: None, + })) + } +} + +fn unique_suffix() -> u128 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() +} + +/// The runtime is returned alongside the session so it outlives the turn and +/// keeps the agent's lease held. +fn session_for(provider: ForcedToolProvider) -> (Runtime, Session) { + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let session = runtime + .create_session("typed-output", model) + .expect("create session"); + (runtime, session) +} + +fn report_spec() -> TerminalOutputSpec { + TerminalOutputSpec::new( + "finish-report", + "Return the final report", + json!({ + "type": "object", + "properties": { + "answer": { "type": "integer" }, + "evidence": { "type": "array", "items": { "type": "string" } } + }, + "required": ["answer", "evidence"] + }), + ) +} + +fn drain(rx: &mut crate::session::SessionEventReceiver) -> Vec { + std::iter::from_fn(|| rx.try_recv().ok()).collect() +} + +fn position( + events: &[SessionEvent], + label: &str, + predicate: impl Fn(&SessionEvent) -> bool, +) -> usize { + events + .iter() + .position(predicate) + .unwrap_or_else(|| panic!("expected a {label} event, got: {events:?}")) +} + +#[tokio::test] +async fn a_typed_turn_returns_the_value_and_counts_as_a_turn() { + let (_runtime, mut session) = session_for(ForcedToolProvider::answering( + json!({ "answer": 42, "evidence": ["a", "b"] }), + )); + + let output = session + .append_turn_to_output::( + vec![ContentBlock::text("produce the report")], + RunOptions::default(), + report_spec(), + ) + .await + .expect("a typed turn succeeds"); + + assert_eq!( + output.value, + Report { + answer: 42, + evidence: vec!["a".to_string(), "b".to_string()], + } + ); + // The turn ends on the tool-result message, not on assistant text — the + // asymmetry the session's terminal event has to account for. + assert_eq!(output.message.role, Role::User); + assert!(matches!( + output.message.content.as_slice(), + [ContentBlock::ToolResult { tool_use_id, .. }] if tool_use_id == "terminal-call-0" + )); + + assert_eq!(session.metadata().turn_count, 1); + assert_eq!(session.metadata().status, SessionStatus::Idle); + assert!( + !session.replay().items().is_empty(), + "the typed turn is committed to the transcript" + ); +} + +#[tokio::test] +async fn a_typed_turn_completes_with_the_model_prose_after_the_terminal_tool_events() { + let (_runtime, mut session) = session_for( + ForcedToolProvider::answering(json!({ "answer": 7, "evidence": [] })) + .with_preface("here is the report"), + ); + let mut rx = session.subscribe(); + + session + .append_turn_to_output::( + vec![ContentBlock::text("produce the report")], + RunOptions::default(), + report_spec(), + ) + .await + .expect("a typed turn succeeds"); + + let events = drain(&mut rx); + + let user = position( + &events, + "UserMessage", + |event| matches!(event, SessionEvent::UserMessage { text } if text == "produce the report"), + ); + let queued = position(&events, "ToolQueued", |event| { + matches!(event, SessionEvent::ToolQueued { tool_name, .. } + if tool_name.starts_with("mentra_terminal_")) + }); + let started = position(&events, "ToolStarted", |event| { + matches!(event, SessionEvent::ToolStarted { .. }) + }); + let completed = position(&events, "ToolCompleted", |event| { + matches!( + event, + SessionEvent::ToolCompleted { + is_error: false, + .. + } + ) + }); + let done = position(&events, "AssistantMessageCompleted", |event| { + matches!(event, SessionEvent::AssistantMessageCompleted { .. }) + }); + + assert!( + user < queued && queued < started && started < completed && completed < done, + "a typed turn runs user -> terminal tool -> completion, got: {events:?}" + ); + + // The completion carries what the model wrote, matching the deltas that + // were already streamed — not the typed payload. + assert!( + matches!(&events[done], SessionEvent::AssistantMessageCompleted { text } + if text == "here is the report"), + "got: {:?}", + events[done] + ); + assert_eq!( + events + .iter() + .filter(|event| matches!(event, SessionEvent::AssistantMessageCompleted { .. })) + .count(), + 1, + "one completion per turn, as for any other turn" + ); + assert!( + events.iter().any(|event| matches!( + event, + SessionEvent::AssistantTokenDelta { full_text, .. } if full_text == "here is the report" + )), + "the completion agrees with the streamed deltas, got: {events:?}" + ); + assert!( + !events + .iter() + .any(|event| matches!(event, SessionEvent::Error { .. })), + "a successful typed turn reports no error, got: {events:?}" + ); + + // The payload is on the stream through the terminal tool's own events, + // which is why the completion does not repeat it as prose. + let SessionEvent::ToolQueued { input_json, .. } = &events[queued] else { + panic!("expected ToolQueued at {queued}"); + }; + assert_eq!( + serde_json::from_str::(input_json).expect("the queued input is JSON"), + json!({ "answer": 7, "evidence": [] }) + ); +} + +#[tokio::test] +async fn a_typed_turn_without_model_prose_completes_with_empty_text() { + let (_runtime, mut session) = session_for(ForcedToolProvider::answering( + json!({ "answer": 1, "evidence": [] }), + )); + let mut rx = session.subscribe(); + + session + .append_turn_to_output::( + vec![ContentBlock::text("produce the report")], + RunOptions::default(), + report_spec(), + ) + .await + .expect("a typed turn succeeds"); + + let events = drain(&mut rx); + let done = position(&events, "AssistantMessageCompleted", |event| { + matches!(event, SessionEvent::AssistantMessageCompleted { .. }) + }); + assert!( + matches!(&events[done], SessionEvent::AssistantMessageCompleted { text } if text.is_empty()), + "a model that writes only the terminal call completes with no prose, got: {:?}", + events[done] + ); +} + +#[tokio::test] +async fn a_value_that_does_not_match_the_type_fails_the_turn_and_the_session_recovers() { + let (_runtime, mut session) = session_for(ForcedToolProvider::answering( + json!({ "answer": "forty-two", "evidence": [] }), + )); + let mut rx = session.subscribe(); + + let error = session + .append_turn_to_output::( + vec![ContentBlock::text("produce the report")], + RunOptions::default(), + report_spec(), + ) + .await + .expect_err("a value that is not a Report must fail the turn"); + + assert!( + error + .to_string() + .contains("did not match the requested type"), + "got: {error}" + ); + assert!( + matches!(session.metadata().status, SessionStatus::Failed(_)), + "a failed typed turn leaves the session Failed, like any failed turn" + ); + assert_eq!( + session.metadata().turn_count, + 0, + "a failed turn does not move the counter" + ); + + let events = drain(&mut rx); + assert!( + events.iter().any(|event| matches!( + event, + SessionEvent::Error { + recoverable: false, + .. + } + )), + "expected a terminal Error event, got: {events:?}" + ); + assert!( + !events + .iter() + .any(|event| matches!(event, SessionEvent::AssistantMessageCompleted { .. })), + "a failed turn emits no completion, got: {events:?}" + ); + + // Same as a failed `append_turn`: the session takes the next turn. + let recovered = session + .append_turn(vec![ContentBlock::text("try again")]) + .await + .expect("the session accepts a turn after a failed typed turn"); + assert_eq!(recovered.text(), PLAIN_ANSWER); + assert_eq!(session.metadata().status, SessionStatus::Idle); + assert_eq!(session.metadata().turn_count, 1); +} + +#[tokio::test] +async fn a_run_that_never_calls_the_terminal_tool_fails_the_turn() { + let (_runtime, mut session) = session_for(ForcedToolProvider::ignoring_the_forced_tool()); + let mut rx = session.subscribe(); + + let error = session + .append_turn_to_output::( + vec![ContentBlock::text("produce the report")], + RunOptions::default(), + report_spec(), + ) + .await + .expect_err("a run without the terminal call has no typed value to return"); + + assert!( + error + .to_string() + .contains("without invoking the expected terminal tool"), + "got: {error}" + ); + assert!(matches!( + session.metadata().status, + SessionStatus::Failed(_) + )); + assert_eq!(session.metadata().turn_count, 0); + + let events = drain(&mut rx); + assert!( + events + .iter() + .any(|event| matches!(event, SessionEvent::Error { .. })), + "expected an Error event, got: {events:?}" + ); +} + +/// A tool an ordinary turn would hold, for the working typed turn below to +/// reach. +struct LookupTool; + +impl ToolDefinition for LookupTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("lookup") + .description("test tool: returns a fact the report needs") + .input_schema(json!({ "type": "object", "properties": {} })) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .build() + } +} + +#[async_trait] +impl ToolExecutor for LookupTool { + async fn execute_mut(&self, _ctx: ToolContext<'_>, _input: Value) -> ToolResult { + Ok("the answer is 42".to_string()) + } +} + +/// Plays a model on a *working* typed turn: it looks the terminal tool up by +/// name in the request rather than being told which to call, works one round, +/// then answers. Its rounds are counted rather than scripted because the two +/// differ only in what the model decides to do. +#[derive(Clone)] +struct WorkingProvider { + model: ModelInfo, + calls: Arc, +} + +impl WorkingProvider { + fn new() -> Self { + Self { + model: ModelInfo::new("typed-output-model", BuiltinProvider::Anthropic), + calls: Arc::new(AtomicUsize::new(0)), + } + } +} + +#[async_trait] +impl Provider for WorkingProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.model.provider.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(vec![self.model.clone()]) + } + + async fn stream(&self, request: Request<'_>) -> Result { + let call = self.calls.fetch_add(1, Ordering::SeqCst); + let terminal = request + .tools + .iter() + .find(|tool| tool.name.starts_with("mentra_terminal_")) + .map(|tool| tool.name.clone()) + .expect("a typed turn always offers its terminal tool"); + assert!( + request.tools.iter().any(|tool| tool.name == "lookup"), + "a working typed turn keeps its ordinary tools: {:?}", + request.tools + ); + assert_eq!( + request.tool_choice, + Some(ToolChoice::Auto), + "and forces none of them" + ); + + let (id, name, input) = if call == 0 { + ("lookup-call".to_string(), "lookup".to_string(), json!({})) + } else { + ( + "terminal-call-0".to_string(), + terminal, + json!({ "answer": 42, "evidence": ["looked it up"] }), + ) + }; + + Ok(provider_event_stream_from_response(Response { + id: format!("message-{call}-{}", unique_suffix()), + model: self.model.id.clone(), + role: Role::Assistant, + content: vec![ContentBlock::ToolUse { id, name, input }], + stop_reason: Some("tool_use".to_string()), + usage: None, + })) + } +} + +#[tokio::test] +async fn a_working_typed_turn_puts_its_tool_work_on_the_session_stream() { + // The reason to want this mode through a `Session` rather than a bare + // agent: the work it does on the way to the answer reaches the same + // stream every other turn's work does, in order, and the typed value + // still comes back. + let provider = WorkingProvider::new(); + let model = provider.model.clone(); + let runtime = Runtime::empty_builder() + .with_provider_instance(provider) + .with_tool(LookupTool) + .build() + .expect("build runtime"); + let mut session = runtime + .create_session("typed-output", model) + .expect("create session"); + let mut rx = session.subscribe(); + + let output = session + .append_turn_to_output::( + vec![ContentBlock::text("look it up, then report")], + RunOptions::default(), + report_spec().with_tools(), + ) + .await + .expect("a working typed turn answers"); + + assert_eq!( + output.value, + Report { + answer: 42, + evidence: vec!["looked it up".to_string()], + } + ); + assert_eq!(session.metadata().turn_count, 1); + assert_eq!(session.metadata().status, SessionStatus::Idle); + + let events = drain(&mut rx); + let looked_up = position( + &events, + "ToolQueued for lookup", + |event| matches!(event, SessionEvent::ToolQueued { tool_name, .. } if tool_name == "lookup"), + ); + let answered = position(&events, "ToolQueued for the terminal tool", |event| { + matches!(event, SessionEvent::ToolQueued { tool_name, .. } + if tool_name.starts_with("mentra_terminal_")) + }); + assert!( + looked_up < answered, + "the turn worked before it answered, got: {events:?}" + ); + assert!( + !events + .iter() + .any(|event| matches!(event, SessionEvent::Error { .. })), + "a successful working typed turn reports no error, got: {events:?}" + ); +} diff --git a/vendor/mentra/src/session/types.rs b/vendor/mentra/src/session/types.rs new file mode 100644 index 0000000..0e408b2 --- /dev/null +++ b/vendor/mentra/src/session/types.rs @@ -0,0 +1,94 @@ +use std::fmt; +use std::str::FromStr; + +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +pub struct SessionId(String); + +impl SessionId { + pub fn new() -> Self { + use rand::Rng; + let nonce: u64 = rand::rng().random(); + Self(format!( + "session-{:x}-{:x}", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(), + nonce + )) + } + + pub fn from_raw(raw: impl Into) -> Self { + Self(raw.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl fmt::Display for SessionId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.0) + } +} + +impl FromStr for SessionId { + type Err = std::convert::Infallible; + + fn from_str(s: &str) -> Result { + Ok(Self(s.to_string())) + } +} + +impl Default for SessionId { + fn default() -> Self { + Self::new() + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum SessionStatus { + #[default] + Created, + Active, + Idle, + Compacting, + Failed(String), +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct SessionMetadata { + pub id: SessionId, + pub title: String, + pub model: String, + pub status: SessionStatus, + pub turn_count: usize, + pub created_at: u64, + pub updated_at: u64, +} + +impl SessionMetadata { + pub fn new(id: SessionId, title: impl Into, model: impl Into) -> Self { + let now = unix_now(); + Self { + id, + title: title.into(), + model: model.into(), + status: SessionStatus::Created, + turn_count: 0, + created_at: now, + updated_at: now, + } + } +} + +fn unix_now() -> u64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} diff --git a/vendor/mentra/src/team.rs b/vendor/mentra/src/team.rs new file mode 100644 index 0000000..fb97de9 --- /dev/null +++ b/vendor/mentra/src/team.rs @@ -0,0 +1,23 @@ +mod actor; +mod host; +mod intrinsic; +mod manager; +mod observer; +mod prompt; +mod store; +mod types; + +pub(crate) use actor::teammate_actor_loop; +pub(crate) use host::{TeammateActorHandle, TeammateHost}; +pub(crate) use manager::TeamManager; +pub(crate) use observer::{TeamObserverSink, TeamRegistration}; +pub(crate) use prompt::{TEAMMATE_MAX_ROUNDS, build_teammate_system_prompt}; +pub(crate) use store::TeamStore; +pub(crate) use types::format_inbox; +pub use types::{ + TeamDispatch, TeamMemberStatus, TeamMemberSummary, TeamMessage, TeamMessageKind, + TeamProtocolRequestSummary, TeamProtocolStatus, +}; +pub(crate) use types::{TeamRequestDirection, TeamRequestFilter}; + +pub(crate) use intrinsic::TeamIntrinsicTool; diff --git a/vendor/mentra/src/team/actor.rs b/vendor/mentra/src/team/actor.rs new file mode 100644 index 0000000..8d03890 --- /dev/null +++ b/vendor/mentra/src/team/actor.rs @@ -0,0 +1,293 @@ +use std::{ + path::{Path, PathBuf}, + sync::Arc, + time::Instant, +}; + +use tokio::sync::{Mutex as AsyncMutex, mpsc}; + +use crate::{ + Agent, ContentBlock, agent::TeamAutonomyConfig, error::RuntimeError, runtime::CancellationToken, +}; + +use super::{TeamManager, TeamMemberStatus}; + +const TEAM_WAKE_PROMPT: &str = "Process any new team inbox messages and continue your work."; +const BACKGROUND_WAKE_PROMPT: &str = + "Review any completed background task results and continue your work."; + +pub(crate) async fn teammate_actor_loop( + manager: TeamManager, + team_dir: PathBuf, + teammate_name: String, + agent: Arc>, + mut wake_rx: mpsc::UnboundedReceiver<()>, + cancellation: CancellationToken, +) { + let autonomy = { + let guard = agent.lock().await; + guard.config().team.autonomy.clone() + }; + let mut should_process = false; + let mut idle_since = None; + + loop { + if cancellation.is_cancelled() { + break; + } + + if should_process { + match process_pending_work(&manager, &team_dir, &teammate_name, &agent).await { + Ok(ActorState::Idle) => { + idle_since.get_or_insert_with(Instant::now); + } + Ok(ActorState::Shutdown) => break, + Err(()) => { + idle_since.get_or_insert_with(Instant::now); + } + } + should_process = false; + } + + if cancellation.is_cancelled() { + break; + } + + if autonomy.enabled { + let started_idle_at = idle_since.unwrap_or_else(|| { + let now = Instant::now(); + idle_since = Some(now); + now + }); + + let wait = tokio::time::sleep(autonomy.poll_interval); + tokio::pin!(wait); + + tokio::select! { + wake = wake_rx.recv() => { + match wake { + Some(()) => { + should_process = true; + idle_since = None; + } + None => break, + } + } + _ = &mut wait => { + match autonomy_tick( + &manager, + &team_dir, + &teammate_name, + &agent, + started_idle_at, + &autonomy, + ).await { + Ok(AutonomyState::ContinueIdle) => {} + Ok(AutonomyState::Claimed(prompt)) => { + if execute_prompt(&manager, &team_dir, &teammate_name, &agent, prompt) + .await + .is_err() + { + idle_since.get_or_insert_with(Instant::now); + continue; + } + let _ = manager.update_member_status( + &team_dir, + &teammate_name, + TeamMemberStatus::Idle, + ); + should_process = true; + idle_since = Some(Instant::now()); + } + Ok(AutonomyState::Shutdown) => break, + Err(()) => { + idle_since.get_or_insert_with(Instant::now); + } + } + } + } + } else { + match wake_rx.recv().await { + Some(()) => { + should_process = true; + idle_since = None; + } + None => break, + } + } + } + + let _ = manager.unregister_teammate_actor(&team_dir, &teammate_name); +} + +enum ActorState { + Idle, + Shutdown, +} + +enum PendingWork { + Prompt(String), + Idle, + Shutdown, +} + +enum AutonomyState { + ContinueIdle, + Claimed(String), + Shutdown, +} + +async fn process_pending_work( + manager: &TeamManager, + team_dir: &Path, + teammate_name: &str, + agent: &Arc>, +) -> Result { + let mut processed_prompt = false; + + loop { + match next_pending_work(manager, team_dir, teammate_name, agent).await { + Ok(PendingWork::Prompt(prompt)) => { + processed_prompt = true; + execute_prompt(manager, team_dir, teammate_name, agent, prompt).await? + } + Ok(PendingWork::Idle) => { + if processed_prompt { + let _ = manager.update_member_status( + team_dir, + teammate_name, + TeamMemberStatus::Idle, + ); + } + return Ok(ActorState::Idle); + } + Ok(PendingWork::Shutdown) => { + let _ = manager.update_member_status( + team_dir, + teammate_name, + TeamMemberStatus::Shutdown, + ); + return Ok(ActorState::Shutdown); + } + Err(error) => { + let _ = mark_failed(manager, team_dir, teammate_name, error); + return Err(()); + } + } + } +} + +async fn next_pending_work( + manager: &TeamManager, + team_dir: &Path, + teammate_name: &str, + agent: &Arc>, +) -> Result { + if manager.has_pending_messages(team_dir, teammate_name)? { + return Ok(PendingWork::Prompt(TEAM_WAKE_PROMPT.to_string())); + } + + let has_background_notifications = { + let guard = agent.lock().await; + guard + .runtime_handle() + .has_deliverable_background_notifications(guard.id()) + }; + if has_background_notifications { + return Ok(PendingWork::Prompt(BACKGROUND_WAKE_PROMPT.to_string())); + } + + if manager.take_shutdown_signal(team_dir, teammate_name)? { + return Ok(PendingWork::Shutdown); + } + + Ok(PendingWork::Idle) +} + +async fn autonomy_tick( + manager: &TeamManager, + team_dir: &Path, + teammate_name: &str, + agent: &Arc>, + idle_since: Instant, + autonomy: &TeamAutonomyConfig, +) -> Result { + match manager.take_shutdown_signal(team_dir, teammate_name) { + Ok(true) => { + let _ = + manager.update_member_status(team_dir, teammate_name, TeamMemberStatus::Shutdown); + return Ok(AutonomyState::Shutdown); + } + Ok(false) => {} + Err(error) => { + let _ = mark_failed(manager, team_dir, teammate_name, error); + return Err(()); + } + } + + let claimed = { + let mut guard = agent.lock().await; + match guard.try_claim_ready_task() { + Ok(task) => task, + Err(error) => { + let _ = mark_failed(manager, team_dir, teammate_name, error); + return Err(()); + } + } + }; + if let Some(task) = claimed { + let task_body = if task.description.trim().is_empty() { + format!("Task #{}: {}", task.id, task.subject) + } else { + format!( + "Task #{}: {}\nDescription: {}", + task.id, task.subject, task.description + ) + }; + return Ok(AutonomyState::Claimed(format!( + "{task_body}\nUpdate your task status. Mark it in_progress when you start and completed when you finish." + ))); + } + + if idle_since.elapsed() >= autonomy.idle_timeout { + let _ = manager.update_member_status(team_dir, teammate_name, TeamMemberStatus::Shutdown); + return Ok(AutonomyState::Shutdown); + } + + Ok(AutonomyState::ContinueIdle) +} + +async fn execute_prompt( + manager: &TeamManager, + team_dir: &Path, + teammate_name: &str, + agent: &Arc>, + prompt: String, +) -> Result<(), ()> { + let _ = manager.update_member_status(team_dir, teammate_name, TeamMemberStatus::Working); + let result = { + let mut guard = agent.lock().await; + guard.send(vec![ContentBlock::Text { text: prompt }]).await + }; + + match result { + Ok(_) | Err(RuntimeError::EmptyAssistantResponse) => Ok(()), + Err(error) => { + let _ = mark_failed(manager, team_dir, teammate_name, error); + Err(()) + } + } +} + +fn mark_failed( + manager: &TeamManager, + team_dir: &Path, + teammate_name: &str, + error: RuntimeError, +) -> Result<(), RuntimeError> { + manager.update_member_status( + team_dir, + teammate_name, + TeamMemberStatus::Failed(error.to_string()), + ) +} diff --git a/vendor/mentra/src/team/host.rs b/vendor/mentra/src/team/host.rs new file mode 100644 index 0000000..bdc3dc9 --- /dev/null +++ b/vendor/mentra/src/team/host.rs @@ -0,0 +1,91 @@ +use std::{path::PathBuf, sync::Arc}; + +use tokio::{ + runtime::{Builder as RuntimeBuilder, Handle as TokioHandle, Runtime as TokioRuntime}, + sync::{Mutex as AsyncMutex, mpsc}, + task::{AbortHandle, JoinHandle}, +}; + +use crate::{agent::Agent, error::RuntimeError, runtime::CancellationToken}; + +use super::{TeamManager, teammate_actor_loop}; + +#[derive(Clone)] +pub(crate) struct TeammateHost { + backend: Arc, +} + +enum TeammateRuntimeBackend { + Current(TokioHandle), + Owned(Arc), +} + +pub(crate) struct TeammateActorHandle { + pub(crate) wake_tx: mpsc::UnboundedSender<()>, + pub(crate) cancellation: CancellationToken, + pub(crate) abort: AbortHandle, +} + +impl Drop for TeammateActorHandle { + fn drop(&mut self) { + self.cancellation.cancel(); + self.abort.abort(); + } +} + +impl TeammateHost { + pub(crate) fn new() -> Result { + let backend = match TokioHandle::try_current() { + Ok(handle) => TeammateRuntimeBackend::Current(handle), + Err(_) => TeammateRuntimeBackend::Owned(Arc::new( + RuntimeBuilder::new_multi_thread() + .worker_threads(2) + .enable_all() + .build() + .map_err(|error| { + RuntimeError::Store(format!( + "Failed to create shared teammate runtime: {error}" + )) + })?, + )), + }; + Ok(Self { + backend: Arc::new(backend), + }) + } + + pub(crate) fn spawn_teammate( + &self, + manager: TeamManager, + team_dir: PathBuf, + teammate_name: String, + agent: Arc>, + ) -> TeammateActorHandle { + let (wake_tx, wake_rx) = mpsc::unbounded_channel(); + let cancellation = CancellationToken::default(); + let task = self.spawn_task(teammate_actor_loop( + manager, + team_dir, + teammate_name, + agent, + wake_rx, + cancellation.clone(), + )); + let abort = task.abort_handle(); + TeammateActorHandle { + wake_tx, + cancellation, + abort, + } + } + + fn spawn_task( + &self, + future: impl std::future::Future + Send + 'static, + ) -> JoinHandle<()> { + match self.backend.as_ref() { + TeammateRuntimeBackend::Current(handle) => handle.spawn(future), + TeammateRuntimeBackend::Owned(runtime) => runtime.handle().spawn(future), + } + } +} diff --git a/vendor/mentra/src/team/intrinsic.rs b/vendor/mentra/src/team/intrinsic.rs new file mode 100644 index 0000000..72394fd --- /dev/null +++ b/vendor/mentra/src/team/intrinsic.rs @@ -0,0 +1,72 @@ +mod execute; +mod schema; + +use async_trait::async_trait; +use strum::{Display, VariantArray}; + +use crate::{ + ContentBlock, + tool::{ + ParallelToolContext, RuntimeToolDescriptor, ToolCall, ToolContext, ToolDefinition, + ToolExecutor, ToolResult, internal::content_block_to_tool_result, + }, +}; + +#[derive(Clone, Copy, Debug, Display, VariantArray)] +#[strum(prefix = "team_")] +#[strum(serialize_all = "snake_case")] +pub(crate) enum TeamIntrinsicTool { + Spawn, + Send, + ReadInbox, + Broadcast, + Request, + Respond, + ListRequests, +} + +impl ToolDefinition for TeamIntrinsicTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + self.tool_spec() + } +} + +#[async_trait] +impl ToolExecutor for TeamIntrinsicTool { + async fn execute(&self, ctx: ParallelToolContext, input: serde_json::Value) -> ToolResult { + match self { + Self::ListRequests => execute::execute_team_list_requests_parallel(ctx, input), + _ => Err(format!( + "Tool '{}' does not support parallel execution", + self.descriptor().provider.name + )), + } + } + + async fn execute_mut(&self, ctx: ToolContext<'_>, input: serde_json::Value) -> ToolResult { + match self { + Self::ListRequests => execute::execute_team_list_requests_parallel(ctx.into(), input), + _ => { + let call = ToolCall { + id: ctx.tool_call_id.clone(), + name: self.to_string(), + input, + }; + let block = match self { + Self::Spawn => execute::execute_team_spawn(ctx.agent, call).await, + Self::Send => execute::execute_team_send(ctx.agent, call), + Self::ReadInbox => execute::execute_team_read_inbox(ctx.agent, call), + Self::Broadcast => execute::execute_team_broadcast(ctx.agent, call), + Self::Request => execute::execute_team_request(ctx.agent, call), + Self::Respond => execute::execute_team_respond(ctx.agent, call), + Self::ListRequests => unreachable!("handled above"), + }; + content_block_to_result(block) + } + } + } +} + +fn content_block_to_result(block: ContentBlock) -> ToolResult { + content_block_to_tool_result("Team intrinsic", block) +} diff --git a/vendor/mentra/src/team/intrinsic/execute.rs b/vendor/mentra/src/team/intrinsic/execute.rs new file mode 100644 index 0000000..8f087a5 --- /dev/null +++ b/vendor/mentra/src/team/intrinsic/execute.rs @@ -0,0 +1,231 @@ +use crate::{ + ContentBlock, + agent::Agent, + error::RuntimeError, + team::{TeamProtocolStatus, TeamRequestDirection, TeamRequestFilter}, + tool::{ParallelToolContext, ToolCall, ToolResult}, +}; + +use super::schema::{ + TeamBroadcastInput, TeamListRequestsInput, TeamRequestInput, TeamRespondInput, TeamSendInput, + TeamSpawnInput, +}; + +pub(super) async fn execute_team_spawn(agent: &mut Agent, call: ToolCall) -> ContentBlock { + let input = match serde_json::from_value::(call.input) { + Ok(input) => input, + Err(error) => { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Invalid team_spawn input: {error}").into(), + is_error: true, + }; + } + }; + + match agent + .spawn_teammate(input.name, input.role, input.prompt) + .await + { + Ok(teammate) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!( + "Spawned persistent teammate '{}' (role: {}, status: {})", + teammate.name, teammate.role, teammate.status + ) + .into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to spawn teammate: {error}").into(), + is_error: true, + }, + } +} + +pub(super) fn execute_team_send(agent: &mut Agent, call: ToolCall) -> ContentBlock { + let input = match serde_json::from_value::(call.input) { + Ok(input) => input, + Err(error) => { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Invalid team_send input: {error}").into(), + is_error: true, + }; + } + }; + + match agent.send_team_message(&input.to, input.content) { + Ok(dispatch) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Sent message to '{}'", dispatch.teammate).into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to send team message: {error}").into(), + is_error: true, + }, + } +} + +pub(super) fn execute_team_read_inbox(agent: &mut Agent, call: ToolCall) -> ContentBlock { + match agent.read_team_inbox().and_then(|messages| { + serde_json::to_string_pretty(&messages).map_err(RuntimeError::FailedToSerializeTeam) + }) { + Ok(content) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: content.into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to read team inbox: {error}").into(), + is_error: true, + }, + } +} + +pub(super) fn execute_team_broadcast(agent: &mut Agent, call: ToolCall) -> ContentBlock { + let input = match serde_json::from_value::(call.input) { + Ok(input) => input, + Err(error) => { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Invalid broadcast input: {error}").into(), + is_error: true, + }; + } + }; + + match agent.broadcast_team_message(input.content) { + Ok(dispatches) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!( + "Broadcast message sent to {} recipient(s): {}", + dispatches.len(), + dispatches + .into_iter() + .map(|dispatch| dispatch.teammate) + .collect::>() + .join(", ") + ) + .into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to broadcast team message: {error}").into(), + is_error: true, + }, + } +} + +pub(super) fn execute_team_request(agent: &mut Agent, call: ToolCall) -> ContentBlock { + let input = match serde_json::from_value::(call.input) { + Ok(input) => input, + Err(error) => { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Invalid team_request input: {error}").into(), + is_error: true, + }; + } + }; + + match agent.request_team_protocol(&input.to, input.protocol, input.content) { + Ok(request) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!( + "Created team request '{}' for '{}' using protocol '{}'", + request.request_id, request.to, request.protocol + ) + .into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to create team request: {error}").into(), + is_error: true, + }, + } +} + +pub(super) fn execute_team_respond(agent: &mut Agent, call: ToolCall) -> ContentBlock { + let input = match serde_json::from_value::(call.input) { + Ok(input) => input, + Err(error) => { + return ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Invalid team_respond input: {error}").into(), + is_error: true, + }; + } + }; + + match agent.respond_team_protocol(&input.request_id, input.approve, input.reason) { + Ok(request) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!( + "{} team request '{}' ({})", + if input.approve { + "Approved" + } else { + "Rejected" + }, + request.request_id, + request.protocol + ) + .into(), + is_error: false, + }, + Err(error) => ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Failed to respond to team request: {error}").into(), + is_error: true, + }, + } +} + +pub(super) fn execute_team_list_requests_parallel( + ctx: ParallelToolContext, + input: serde_json::Value, +) -> ToolResult { + let filter = parse_team_request_filter(input)?; + let config = ctx.runtime.agent_config(&ctx.agent_id)?; + let requests = ctx + .runtime + .list_team_requests(&config.team_dir, &config.name, filter) + .map_err(|error| format!("Failed to list team requests: {error}"))?; + + serde_json::to_string_pretty(&requests) + .map_err(|error| format!("Failed to serialize team requests: {error}")) +} + +fn parse_team_request_filter(input: serde_json::Value) -> Result { + let input = serde_json::from_value::(input) + .map_err(|error| format!("Invalid team_list_requests input: {error}"))?; + + let status = match input.status.as_deref() { + Some("pending") => Some(TeamProtocolStatus::Pending), + Some("approved") => Some(TeamProtocolStatus::Approved), + Some("rejected") => Some(TeamProtocolStatus::Rejected), + Some(value) => return Err(format!("Invalid team_list_requests status '{value}'")), + None => None, + }; + + let direction = match input.direction.as_deref() { + Some("inbound") => TeamRequestDirection::Inbound, + Some("outbound") => TeamRequestDirection::Outbound, + Some("any") | None => TeamRequestDirection::Any, + Some(value) => return Err(format!("Invalid team_list_requests direction '{value}'")), + }; + + Ok(TeamRequestFilter { + status, + protocol: input.protocol, + counterparty: input.counterparty, + direction, + }) +} diff --git a/vendor/mentra/src/team/intrinsic/schema.rs b/vendor/mentra/src/team/intrinsic/schema.rs new file mode 100644 index 0000000..39287a5 --- /dev/null +++ b/vendor/mentra/src/team/intrinsic/schema.rs @@ -0,0 +1,208 @@ +use serde::Deserialize; +use serde_json::json; + +use crate::tool::{ + RuntimeToolDescriptor, ToolApprovalCategory, ToolCapability, ToolDurability, + ToolExecutionCategory, ToolSideEffectLevel, + internal::{RuntimeDescriptorParts, build_runtime_descriptor}, +}; + +use super::TeamIntrinsicTool; + +#[derive(Debug, Deserialize)] +pub(super) struct TeamSpawnInput { + pub(super) name: String, + pub(super) role: String, + pub(super) prompt: Option, +} + +#[derive(Debug, Deserialize)] +pub(super) struct TeamSendInput { + pub(super) to: String, + pub(super) content: String, +} + +#[derive(Debug, Deserialize)] +pub(super) struct TeamBroadcastInput { + pub(super) content: String, +} + +#[derive(Debug, Deserialize)] +pub(super) struct TeamRequestInput { + pub(super) to: String, + pub(super) protocol: String, + pub(super) content: String, +} + +#[derive(Debug, Deserialize)] +pub(super) struct TeamRespondInput { + pub(super) request_id: String, + pub(super) approve: bool, + pub(super) reason: Option, +} + +#[derive(Debug, Deserialize)] +pub(super) struct TeamListRequestsInput { + pub(super) status: Option, + pub(super) protocol: Option, + pub(super) counterparty: Option, + pub(super) direction: Option, +} + +impl TeamIntrinsicTool { + fn team_spec( + &self, + description: &str, + input_schema: serde_json::Value, + execution_category: ToolExecutionCategory, + ) -> RuntimeToolDescriptor { + build_runtime_descriptor(RuntimeDescriptorParts { + name: self.to_string(), + description: description.to_string(), + input_schema, + capabilities: vec![ToolCapability::TeamCoordination], + side_effect_level: ToolSideEffectLevel::LocalState, + durability: ToolDurability::Persistent, + execution_category, + approval_category: ToolApprovalCategory::Delegation, + }) + } + + pub(super) fn tool_spec(&self) -> RuntimeToolDescriptor { + match self { + TeamIntrinsicTool::Spawn => self.team_spec( + "Create a persistent teammate that can receive mailbox messages across turns.", + json!({ + "type": "object", + "properties": { + "name": { + "type": "string", + "description": "Unique teammate name" + }, + "role": { + "type": "string", + "description": "Short responsibility or specialty for this teammate" + }, + "prompt": { + "type": "string", + "description": "Optional kickoff message to send immediately after spawning" + } + }, + "required": ["name", "role"] + }), + ToolExecutionCategory::Delegation, + ), + TeamIntrinsicTool::Send => self.team_spec( + "Send a normal mailbox message to the lead or a persistent teammate. Use this to ask a teammate for work or a proposal; do not use team_request when you are simply asking them to submit a plan back to you.", + json!({ + "type": "object", + "properties": { + "to": { + "type": "string", + "description": "Recipient teammate or lead name" + }, + "content": { + "type": "string", + "description": "Message body to deliver" + } + }, + "required": ["to", "content"] + }), + ToolExecutionCategory::ExclusivePersistentMutation, + ), + TeamIntrinsicTool::ReadInbox => self.team_spec( + "Read and drain any currently pending mailbox messages for this agent.", + json!({ + "type": "object", + "properties": {} + }), + ToolExecutionCategory::ExclusivePersistentMutation, + ), + TeamIntrinsicTool::Broadcast => self.team_spec( + "Lead-only team announcement tool. Send the same mailbox message to every other known agent on the team.", + json!({ + "type": "object", + "properties": { + "content": { + "type": "string", + "description": "Message body to deliver to every other teammate" + } + }, + "required": ["content"] + }), + ToolExecutionCategory::ExclusivePersistentMutation, + ), + TeamIntrinsicTool::Request => self.team_spec( + "Create a structured team request with a generated request_id and durable status. Use this when you are the requester and expect the other side to answer with team_respond. For built-in plan review, the teammate doing risky work should send protocol `plan_approval` to the lead; the lead should usually ask for the plan with team_send, then answer the inbound request with team_respond.", + json!({ + "type": "object", + "properties": { + "to": { + "type": "string", + "description": "Recipient teammate or lead name" + }, + "protocol": { + "type": "string", + "description": "Open-ended protocol kind such as shutdown or plan_approval" + }, + "content": { + "type": "string", + "description": "Request body or plan text" + } + }, + "required": ["to", "protocol", "content"] + }), + ToolExecutionCategory::ExclusivePersistentMutation, + ), + TeamIntrinsicTool::Respond => self.team_spec( + "Approve or reject a pending team request by request_id.", + json!({ + "type": "object", + "properties": { + "request_id": { + "type": "string", + "description": "Correlated request identifier" + }, + "approve": { + "type": "boolean", + "description": "Whether to approve the request" + }, + "reason": { + "type": "string", + "description": "Optional explanation or feedback" + } + }, + "required": ["request_id", "approve"] + }), + ToolExecutionCategory::ExclusivePersistentMutation, + ), + TeamIntrinsicTool::ListRequests => self.team_spec( + "List visible team protocol requests with optional filters.", + json!({ + "type": "object", + "properties": { + "status": { + "type": "string", + "enum": ["pending", "approved", "rejected"], + "description": "Optional request status filter" + }, + "protocol": { + "type": "string", + "description": "Optional protocol kind filter" + }, + "counterparty": { + "type": "string", + "description": "Optional other participant filter" + }, + "direction": { + "type": "string", + "enum": ["inbound", "outbound", "any"], + "description": "Filter relative to the current agent" + } + } + }), + ToolExecutionCategory::ReadOnlyParallel, + ), + } + } +} diff --git a/vendor/mentra/src/team/manager.rs b/vendor/mentra/src/team/manager.rs new file mode 100644 index 0000000..b494cdb --- /dev/null +++ b/vendor/mentra/src/team/manager.rs @@ -0,0 +1,722 @@ +use crate::{agent::AgentEvent, error::RuntimeError}; +use std::{ + collections::{HashMap, HashSet}, + path::{Path, PathBuf}, + sync::{ + Arc, Mutex, + atomic::{AtomicU64, Ordering}, + }, +}; + +use super::{ + TeamDispatch, TeamMemberStatus, TeamMemberSummary, TeamMessage, TeamObserverSink, + TeamProtocolRequestSummary, TeamProtocolStatus, TeamRegistration, TeamRequestFilter, + TeammateActorHandle, +}; +use crate::runtime::RuntimeStore; + +static NEXT_REQUEST_ID: AtomicU64 = AtomicU64::new(1); + +#[derive(Clone)] +pub(crate) struct TeamManager { + inner: Arc, +} + +struct TeamManagerInner { + store: Arc, + state: Mutex, +} + +#[derive(Default)] +struct TeamManagerState { + teams: HashMap, +} + +#[derive(Default)] +struct TeamState { + team_dir: PathBuf, + members: Vec, + requests: Vec, + known_agents: HashSet, + unread_counts: HashMap, + observers: Vec, + actors: HashMap, + pending_shutdowns: HashSet, +} + +#[derive(Clone)] +struct TeamObserver { + agent_name: String, + sink: Arc, +} + +#[derive(Clone)] +struct ObserverUpdate { + observer: TeamObserver, + unread_count: usize, +} + +impl TeamManager { + pub(crate) fn new(store: Arc) -> Self { + Self { + inner: Arc::new(TeamManagerInner { + store, + state: Default::default(), + }), + } + } + + pub(crate) fn register_agent( + &self, + registration: TeamRegistration, + ) -> Result<(), RuntimeError> { + let TeamRegistration { + agent_name, + team_dir, + observer, + } = registration; + let (members, requests, unread_count) = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir.as_path())?; + team.known_agents.insert(agent_name.clone()); + team.unread_counts.insert( + agent_name.clone(), + self.inner + .store + .unread_team_count(team_dir.as_path(), &agent_name)?, + ); + team.observers + .retain(|existing| existing.agent_name != agent_name); + team.observers.push(TeamObserver { + agent_name: agent_name.clone(), + sink: observer.clone(), + }); + ( + team.members.clone(), + team.requests.clone(), + team.unread_counts + .get(&agent_name) + .copied() + .unwrap_or_default(), + ) + }; + + observer.publish_snapshot(&members, &requests, unread_count); + Ok(()) + } + + pub(crate) fn spawn_teammate( + &self, + team_dir: &Path, + summary: TeamMemberSummary, + actor: TeammateActorHandle, + ) -> Result { + let (observer_updates, members, requests) = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + if let Some(index) = team + .members + .iter() + .position(|member| member.name == summary.name) + { + if team.actors.contains_key(&summary.name) { + return Err(RuntimeError::InvalidTeam(format!( + "Team member '{}' already exists", + summary.name + ))); + } + + team.members[index] = summary.clone(); + } else { + team.members.push(summary.clone()); + } + team.known_agents.insert(summary.name.clone()); + team.actors.insert(summary.name.clone(), actor); + self.inner + .store + .upsert_team_member(&team.team_dir, &summary)?; + ( + observer_updates(team), + team.members.clone(), + team.requests.clone(), + ) + }; + + self.publish_to_observers( + observer_updates, + members, + requests, + AgentEvent::TeammateSpawned { + teammate: summary.clone(), + }, + ); + Ok(summary) + } + + pub(crate) fn wake_teammate( + &self, + team_dir: &Path, + teammate_name: &str, + ) -> Result<(), RuntimeError> { + let wake_tx = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + team.actors + .get(teammate_name) + .map(|actor| actor.wake_tx.clone()) + .ok_or_else(|| { + RuntimeError::InvalidTeam(format!( + "No live teammate actor exists for '{teammate_name}'" + )) + })? + }; + + let _ = wake_tx.send(()); + Ok(()) + } + + pub(crate) fn update_member_status( + &self, + team_dir: &Path, + name: &str, + status: TeamMemberStatus, + ) -> Result<(), RuntimeError> { + let (observer_updates, members, requests, teammate) = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + let teammate = team + .members + .iter_mut() + .find(|member| member.name == name) + .ok_or_else(|| { + RuntimeError::InvalidTeam(format!("Unknown team member '{name}'")) + })?; + teammate.status = status; + self.inner + .store + .upsert_team_member(&team.team_dir, teammate)?; + let teammate = teammate.clone(); + ( + observer_updates(team), + team.members.clone(), + team.requests.clone(), + teammate, + ) + }; + + self.publish_to_observers( + observer_updates, + members, + requests, + AgentEvent::TeammateUpdated { teammate }, + ); + Ok(()) + } + + pub(crate) fn send_message( + &self, + team_dir: &Path, + sender: &str, + to: &str, + content: String, + ) -> Result { + let (wake_tx, notification) = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + if !team.known_agents.contains(to) + && !team.members.iter().any(|member| member.name == to) + { + return Err(RuntimeError::InvalidTeam(format!( + "Unknown team recipient '{to}'" + ))); + } + + self.inner.store.append_team_message( + team_dir, + to, + &TeamMessage::message(sender.to_string(), content), + )?; + increment_unread_count(team, to); + + ( + team.actors.get(to).map(|actor| actor.wake_tx.clone()), + inbox_notification(team, to), + ) + }; + + if let Some(wake_tx) = wake_tx { + let _ = wake_tx.send(()); + } + self.publish_inbox_notification(notification); + + Ok(TeamDispatch { + teammate: to.to_string(), + }) + } + + pub(crate) fn broadcast_message( + &self, + team_dir: &Path, + sender: &str, + content: String, + ) -> Result, RuntimeError> { + let (recipients, wake_txs, notifications) = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + + let mut recipients = team.known_agents.iter().cloned().collect::>(); + recipients.sort(); + recipients.retain(|name| name != sender); + + let mut wake_txs = Vec::new(); + for recipient in &recipients { + self.inner.store.append_team_message( + team_dir, + recipient, + &TeamMessage::broadcast(sender.to_string(), content.clone()), + )?; + increment_unread_count(team, recipient); + + if let Some(wake_tx) = team + .actors + .get(recipient) + .map(|actor| actor.wake_tx.clone()) + { + wake_txs.push(wake_tx); + } + } + + let notifications = recipients + .iter() + .map(|recipient| inbox_notification(team, recipient)) + .collect::>(); + + (recipients, wake_txs, notifications) + }; + + for wake_tx in wake_txs { + let _ = wake_tx.send(()); + } + for notification in notifications { + self.publish_inbox_notification(notification); + } + + Ok(recipients + .into_iter() + .map(|teammate| TeamDispatch { teammate }) + .collect()) + } + + pub(crate) fn read_inbox( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result, RuntimeError> { + let (messages, notification) = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + let messages = self.inner.store.read_team_inbox(team_dir, agent_name)?; + team.unread_counts.insert(agent_name.to_string(), 0); + (messages, inbox_notification(team, agent_name)) + }; + self.publish_inbox_notification(notification); + Ok(messages) + } + + pub(crate) fn has_pending_messages( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result { + Ok(self.inner.store.unread_team_count(team_dir, agent_name)? > 0) + } + + pub(crate) fn requeue_messages( + &self, + team_dir: &Path, + agent_name: &str, + _messages: Vec, + ) -> Result<(), RuntimeError> { + let notification = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + self.inner.store.requeue_team_inbox(team_dir, agent_name)?; + team.unread_counts.insert( + agent_name.to_string(), + self.inner.store.unread_team_count(team_dir, agent_name)?, + ); + inbox_notification(team, agent_name) + }; + self.publish_inbox_notification(notification); + Ok(()) + } + + pub(crate) fn acknowledge_messages( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result<(), RuntimeError> { + let notification = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + self.inner.store.ack_team_inbox(team_dir, agent_name)?; + team.unread_counts.insert( + agent_name.to_string(), + self.inner.store.unread_team_count(team_dir, agent_name)?, + ); + inbox_notification(team, agent_name) + }; + self.publish_inbox_notification(notification); + Ok(()) + } + + pub(crate) fn create_request( + &self, + team_dir: &Path, + sender: &str, + to: &str, + protocol: String, + content: String, + ) -> Result { + let (observer_updates, members, requests, request, wake_tx, notification) = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + if !team.known_agents.contains(to) + && !team.members.iter().any(|member| member.name == to) + { + return Err(RuntimeError::InvalidTeam(format!( + "Unknown team recipient '{to}'" + ))); + } + + let request = TeamProtocolRequestSummary { + request_id: next_request_id(), + protocol, + from: sender.to_string(), + to: to.to_string(), + content, + status: TeamProtocolStatus::Pending, + created_at: unix_timestamp_secs(), + resolved_at: None, + resolution_reason: None, + }; + + self.inner.store.append_team_message( + team_dir, + to, + &TeamMessage::request(sender.to_string(), &request), + )?; + increment_unread_count(team, to); + + team.requests.push(request.clone()); + self.inner + .store + .upsert_team_request(&team.team_dir, &request)?; + ( + observer_updates(team), + team.members.clone(), + team.requests.clone(), + request, + team.actors.get(to).map(|actor| actor.wake_tx.clone()), + inbox_notification(team, to), + ) + }; + + if let Some(wake_tx) = wake_tx { + let _ = wake_tx.send(()); + } + self.publish_inbox_notification(notification); + + self.publish_to_observers( + observer_updates, + members, + requests, + AgentEvent::TeamProtocolRequested { + request: request.clone(), + }, + ); + Ok(request) + } + + pub(crate) fn resolve_request( + &self, + team_dir: &Path, + responder: &str, + request_id: &str, + approve: bool, + reason: Option, + ) -> Result { + let (observer_updates, members, requests, request, wake_tx, notification) = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + let request = team + .requests + .iter_mut() + .find(|request| request.request_id == request_id) + .ok_or_else(|| { + RuntimeError::InvalidTeam(format!("Unknown team request '{request_id}'")) + })?; + + if request.to != responder { + return Err(RuntimeError::InvalidTeam(format!( + "Agent '{responder}' cannot respond to request '{request_id}'" + ))); + } + + if request.status != TeamProtocolStatus::Pending { + return Err(RuntimeError::InvalidTeam(format!( + "Team request '{request_id}' is already resolved" + ))); + } + + request.status = if approve { + TeamProtocolStatus::Approved + } else { + TeamProtocolStatus::Rejected + }; + request.resolved_at = Some(unix_timestamp_secs()); + request.resolution_reason = reason.clone().filter(|value| !value.trim().is_empty()); + let request = request.clone(); + let response_body = request.resolution_reason.clone().unwrap_or_default(); + + self.inner.store.append_team_message( + team_dir, + &request.from, + &TeamMessage::response(responder.to_string(), &request, approve, response_body), + )?; + increment_unread_count(team, &request.from); + + if approve && request.protocol == "shutdown" { + team.pending_shutdowns.insert(responder.to_string()); + } + + self.inner + .store + .upsert_team_request(&team.team_dir, &request)?; + let wake_tx = team + .actors + .get(&request.from) + .map(|actor| actor.wake_tx.clone()); + let notification = inbox_notification(team, &request.from); + ( + observer_updates(team), + team.members.clone(), + team.requests.clone(), + request, + wake_tx, + notification, + ) + }; + + if let Some(wake_tx) = wake_tx { + let _ = wake_tx.send(()); + } + self.publish_inbox_notification(notification); + + self.publish_to_observers( + observer_updates, + members, + requests, + AgentEvent::TeamProtocolResolved { + request: request.clone(), + }, + ); + Ok(request) + } + + pub(crate) fn list_requests( + &self, + team_dir: &Path, + agent_name: &str, + filter: TeamRequestFilter, + ) -> Result, RuntimeError> { + let mut requests = { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + team.requests + .iter() + .filter(|request| filter.matches(agent_name, request)) + .cloned() + .collect::>() + }; + requests.sort_by(|left, right| { + left.created_at + .cmp(&right.created_at) + .then_with(|| left.request_id.cmp(&right.request_id)) + }); + Ok(requests) + } + + pub(crate) fn take_shutdown_signal( + &self, + team_dir: &Path, + teammate_name: &str, + ) -> Result { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + Ok(team.pending_shutdowns.remove(teammate_name)) + } + + pub(crate) fn unregister_teammate_actor( + &self, + team_dir: &Path, + teammate_name: &str, + ) -> Result<(), RuntimeError> { + let mut state = self.inner.state.lock().expect("team manager poisoned"); + let team = ensure_team_state(&self.inner.store, &mut state, team_dir)?; + team.actors.remove(teammate_name); + Ok(()) + } + + fn publish_to_observers( + &self, + observer_updates: Vec, + members: Vec, + requests: Vec, + event: AgentEvent, + ) { + for update in observer_updates { + let observer = update.observer; + observer + .sink + .publish_snapshot(&members, &requests, update.unread_count); + observer.sink.publish_event(event.clone()); + } + } + + fn publish_inbox_notification(&self, notification: Option) { + let Some(notification) = notification else { + return; + }; + + for observer in notification.observers { + observer.sink.publish_snapshot( + ¬ification.members, + ¬ification.requests, + notification.unread_count, + ); + observer.sink.publish_event(AgentEvent::TeamInboxUpdated { + unread_count: notification.unread_count, + }); + } + } +} + +fn ensure_team_state<'a>( + store: &Arc, + state: &'a mut TeamManagerState, + team_dir: &Path, +) -> Result<&'a mut TeamState, RuntimeError> { + let key = team_key(team_dir); + if !state.teams.contains_key(&key) { + state.teams.insert( + key.clone(), + TeamState { + team_dir: team_dir.to_path_buf(), + members: store.load_team_members(team_dir)?, + requests: store.load_team_requests(team_dir)?, + known_agents: store.list_team_agent_names(team_dir)?.into_iter().collect(), + ..Default::default() + }, + ); + } + + Ok(state.teams.get_mut(&key).expect("team state missing")) +} + +#[derive(Clone)] +struct InboxNotification { + observers: Vec, + members: Vec, + requests: Vec, + unread_count: usize, +} + +fn team_key(team_dir: &Path) -> String { + team_dir.to_string_lossy().into_owned() +} + +fn increment_unread_count(team: &mut TeamState, agent_name: &str) { + *team + .unread_counts + .entry(agent_name.to_string()) + .or_insert(0) += 1; +} + +fn inbox_notification(team: &TeamState, agent_name: &str) -> Option { + let unread_count = team + .unread_counts + .get(agent_name) + .copied() + .unwrap_or_default(); + let observers = team + .observers + .iter() + .filter(|observer| observer.agent_name == agent_name) + .cloned() + .collect::>(); + if observers.is_empty() { + return None; + } + + Some(InboxNotification { + observers, + members: team.members.clone(), + requests: team.requests.clone(), + unread_count, + }) +} + +fn observer_updates(team: &TeamState) -> Vec { + team.observers + .iter() + .cloned() + .map(|observer| ObserverUpdate { + unread_count: team + .unread_counts + .get(&observer.agent_name) + .copied() + .unwrap_or_default(), + observer, + }) + .collect() +} + +fn next_request_id() -> String { + let counter = NEXT_REQUEST_ID.fetch_add(1, Ordering::Relaxed); + format_request_id(unix_timestamp_nanos(), counter) +} + +fn format_request_id(timestamp_nanos: u128, counter: u64) -> String { + format!("{timestamp_nanos:x}-{counter:x}") +} + +fn unix_timestamp_nanos() -> u128 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() +} + +fn unix_timestamp_secs() -> u64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} + +#[cfg(test)] +mod tests { + use super::format_request_id; + + #[test] + fn request_ids_remain_distinct_for_the_old_cross_second_xor_collision() { + assert_eq!(100_u128 ^ 2, 101_u128 ^ 3); + assert_ne!(format_request_id(100, 2), format_request_id(101, 3)); + } + + #[test] + fn request_ids_remain_distinct_when_only_the_counter_changes() { + assert_ne!(format_request_id(100, 2), format_request_id(100, 3)); + } +} diff --git a/vendor/mentra/src/team/observer.rs b/vendor/mentra/src/team/observer.rs new file mode 100644 index 0000000..ddf0493 --- /dev/null +++ b/vendor/mentra/src/team/observer.rs @@ -0,0 +1,23 @@ +use std::{path::PathBuf, sync::Arc}; + +use crate::agent::AgentEvent; + +use super::{TeamMemberSummary, TeamProtocolRequestSummary}; + +pub(crate) trait TeamObserverSink: Send + Sync { + fn publish_snapshot( + &self, + members: &[TeamMemberSummary], + requests: &[TeamProtocolRequestSummary], + unread_count: usize, + ); + + fn publish_event(&self, event: AgentEvent); +} + +#[derive(Clone)] +pub(crate) struct TeamRegistration { + pub(crate) agent_name: String, + pub(crate) team_dir: PathBuf, + pub(crate) observer: Arc, +} diff --git a/vendor/mentra/src/team/prompt.rs b/vendor/mentra/src/team/prompt.rs new file mode 100644 index 0000000..503a69a --- /dev/null +++ b/vendor/mentra/src/team/prompt.rs @@ -0,0 +1,19 @@ +use std::borrow::Cow; + +pub(crate) const TEAMMATE_MAX_ROUNDS: usize = 50; +const TEAMMATE_SYSTEM_PROMPT: &str = "You are a persistent teammate inside a larger agent team. You may receive new mailbox messages across multiple turns. Use team_send for targeted coordination, team_request to start structured request-response protocols, team_respond to answer protocol requests, and team_list_requests to inspect approval state when needed. For risky or destructive work, wait until the lead asks you for a proposal, then submit your plan with protocol `plan_approval` and wait for the matching response before proceeding. If you receive a shutdown request and decide to approve it, send team_respond, finish your current turn cleanly, and then exit. Finish each turn with a concise progress update."; + +pub(crate) fn build_teammate_system_prompt( + base: Option>, + name: &str, + role: &str, + lead: &str, +) -> String { + let addition = format!( + "You are teammate '{name}' with role '{role}' on a team led by '{lead}'. {TEAMMATE_SYSTEM_PROMPT}" + ); + match base { + Some(system) => format!("{system}\n\n{addition}"), + None => addition, + } +} diff --git a/vendor/mentra/src/team/store.rs b/vendor/mentra/src/team/store.rs new file mode 100644 index 0000000..03780f7 --- /dev/null +++ b/vendor/mentra/src/team/store.rs @@ -0,0 +1,38 @@ +use std::path::Path; + +use crate::error::RuntimeError; + +use super::{TeamMemberSummary, TeamMessage, TeamProtocolRequestSummary}; + +pub trait TeamStore: Send + Sync { + fn unread_team_count(&self, team_dir: &Path, agent_name: &str) -> Result; + fn load_team_members(&self, team_dir: &Path) -> Result, RuntimeError>; + fn upsert_team_member( + &self, + team_dir: &Path, + summary: &TeamMemberSummary, + ) -> Result<(), RuntimeError>; + fn read_team_inbox( + &self, + team_dir: &Path, + agent_name: &str, + ) -> Result, RuntimeError>; + fn ack_team_inbox(&self, team_dir: &Path, agent_name: &str) -> Result<(), RuntimeError>; + fn requeue_team_inbox(&self, team_dir: &Path, agent_name: &str) -> Result<(), RuntimeError>; + fn append_team_message( + &self, + team_dir: &Path, + recipient: &str, + message: &TeamMessage, + ) -> Result<(), RuntimeError>; + fn load_team_requests( + &self, + team_dir: &Path, + ) -> Result, RuntimeError>; + fn upsert_team_request( + &self, + team_dir: &Path, + request: &TeamProtocolRequestSummary, + ) -> Result<(), RuntimeError>; + fn list_team_agent_names(&self, team_dir: &Path) -> Result, RuntimeError>; +} diff --git a/vendor/mentra/src/team/types.rs b/vendor/mentra/src/team/types.rs new file mode 100644 index 0000000..6b7caa9 --- /dev/null +++ b/vendor/mentra/src/team/types.rs @@ -0,0 +1,191 @@ +use std::time::{SystemTime, UNIX_EPOCH}; + +use serde::{Deserialize, Serialize}; +use strum::Display; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default, Display)] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum TeamMemberStatus { + #[default] + Idle, + Working, + #[strum(to_string = "failed: {0}")] + Failed(String), + Shutdown, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TeamMemberSummary { + pub id: String, + pub name: String, + pub role: String, + pub model: String, + pub status: TeamMemberStatus, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default, Display)] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum TeamProtocolStatus { + #[default] + Pending, + Approved, + Rejected, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TeamProtocolRequestSummary { + pub request_id: String, + pub protocol: String, + pub from: String, + pub to: String, + pub content: String, + pub status: TeamProtocolStatus, + pub created_at: u64, + pub resolved_at: Option, + pub resolution_reason: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Display, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum TeamMessageKind { + Message, + Broadcast, + Request, + Response, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TeamMessage { + #[serde(rename = "type")] + pub kind: TeamMessageKind, + #[serde(rename = "from")] + pub sender: String, + pub content: String, + pub timestamp: u64, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub request_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub protocol: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub approve: Option, +} + +impl TeamMessage { + pub(crate) fn message(sender: String, content: String) -> Self { + Self { + kind: TeamMessageKind::Message, + sender, + content, + timestamp: unix_timestamp_secs(), + request_id: None, + protocol: None, + approve: None, + } + } + + pub(crate) fn broadcast(sender: String, content: String) -> Self { + Self { + kind: TeamMessageKind::Broadcast, + sender, + content, + timestamp: unix_timestamp_secs(), + request_id: None, + protocol: None, + approve: None, + } + } + + pub(crate) fn request(sender: String, request: &TeamProtocolRequestSummary) -> Self { + Self { + kind: TeamMessageKind::Request, + sender, + content: request.content.clone(), + timestamp: unix_timestamp_secs(), + request_id: Some(request.request_id.clone()), + protocol: Some(request.protocol.clone()), + approve: None, + } + } + + pub(crate) fn response( + sender: String, + request: &TeamProtocolRequestSummary, + approve: bool, + reason: String, + ) -> Self { + Self { + kind: TeamMessageKind::Response, + sender, + content: reason, + timestamp: unix_timestamp_secs(), + request_id: Some(request.request_id.clone()), + protocol: Some(request.protocol.clone()), + approve: Some(approve), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TeamDispatch { + pub teammate: String, +} + +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub(crate) enum TeamRequestDirection { + Inbound, + Outbound, + #[default] + Any, +} + +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub(crate) struct TeamRequestFilter { + pub status: Option, + pub protocol: Option, + pub counterparty: Option, + pub direction: TeamRequestDirection, +} + +impl TeamRequestFilter { + pub(crate) fn matches(&self, agent_name: &str, request: &TeamProtocolRequestSummary) -> bool { + if let Some(status) = &self.status + && &request.status != status + { + return false; + } + + if let Some(protocol) = &self.protocol + && request.protocol != *protocol + { + return false; + } + + if let Some(counterparty) = &self.counterparty + && request.from != *counterparty + && request.to != *counterparty + { + return false; + } + + match self.direction { + TeamRequestDirection::Inbound => request.to == agent_name, + TeamRequestDirection::Outbound => request.from == agent_name, + TeamRequestDirection::Any => request.from == agent_name || request.to == agent_name, + } + } +} + +pub(crate) fn format_inbox(messages: &[TeamMessage]) -> String { + let body = serde_json::to_string_pretty(messages).unwrap_or_else(|_| "[]".to_string()); + format!("\n{body}\n") +} + +fn unix_timestamp_secs() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} diff --git a/vendor/mentra/src/test.rs b/vendor/mentra/src/test.rs new file mode 100644 index 0000000..90064e1 --- /dev/null +++ b/vendor/mentra/src/test.rs @@ -0,0 +1,715 @@ +use std::{ + collections::VecDeque, + sync::{ + Arc, Mutex, + atomic::{AtomicU64, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use serde_json::Value; +use tokio::sync::mpsc; + +use crate::{ + BuiltinProvider, ModelInfo, Runtime, RuntimePolicy, + error::RuntimeError, + provider::{ + ContentBlock, Provider, ProviderDescriptor, ProviderError, ProviderEvent, + ProviderEventStream, ProviderId, Request, Response, Role, + provider_event_stream_from_response, + }, + runtime::{PreExecutionHook, SqliteRuntimeStore, VolatileRuntimeStore}, + tool::ToolAuthorizer, +}; + +/// Disambiguates mock runtimes built within one clock tick. The wall clock +/// alone is not a source of uniqueness: two builds can read the same +/// nanosecond, and did. +static NEXT_MOCK_RUNTIME_ID: AtomicU64 = AtomicU64::new(0); + +#[derive(Debug)] +pub enum MockTurn { + Text(String), + StreamText(Vec), + ToolCalls(Vec), + Failure(ProviderError), +} + +#[derive(Debug, Clone)] +pub struct MockToolCall { + id: Option, + name: String, + input: Value, +} + +impl MockToolCall { + pub fn new(name: impl Into, input: Value) -> Self { + Self { + id: None, + name: name.into(), + input, + } + } + + pub fn with_id(mut self, id: impl Into) -> Self { + self.id = Some(id.into()); + self + } +} + +/// A runtime whose provider replies from a script, for tests that need a real +/// runtime without a real model. +/// +/// State lives in a [`VolatileRuntimeStore`] unless +/// [`MockRuntimeBuilder::with_store`] says otherwise: a mock writes nothing to +/// disk, and two mocks never share anything. Dropping one leaves no file to +/// clean up. +pub struct MockRuntime { + runtime: Runtime, + provider: ScriptedProvider, + model: ModelInfo, +} + +impl MockRuntime { + pub fn builder() -> MockRuntimeBuilder { + MockRuntimeBuilder::default() + } + + pub fn runtime(&self) -> &Runtime { + &self.runtime + } + + pub fn model(&self) -> ModelInfo { + self.model.clone() + } + + pub async fn recorded_requests(&self) -> Vec> { + self.provider.recorded_requests() + } +} + +pub struct MockRuntimeBuilder { + model: ModelInfo, + turns: Vec, + runtime_identifier: String, + store: Option, + policy: RuntimePolicy, + tool_authorizer: Option>, + pre_hook: Option>, +} + +impl Default for MockRuntimeBuilder { + fn default() -> Self { + let runtime_identifier = format!( + "mock-runtime-{}-{}", + now_nanos(), + NEXT_MOCK_RUNTIME_ID.fetch_add(1, Ordering::Relaxed) + ); + Self { + model: ModelInfo::new("mock-model", BuiltinProvider::OpenAI), + turns: Vec::new(), + runtime_identifier, + store: None, + policy: RuntimePolicy::permissive(), + tool_authorizer: None, + pre_hook: None, + } + } +} + +impl MockRuntimeBuilder { + pub fn model(mut self, id: impl Into, provider: impl Into) -> Self { + self.model = ModelInfo::new(id.into(), provider.into()); + self + } + + pub fn runtime_identifier(mut self, runtime_identifier: impl Into) -> Self { + self.runtime_identifier = runtime_identifier.into(); + self + } + + /// Runs the scripted runtime against a SQLite store instead of the + /// volatile default, for a test that needs state to outlive the + /// `MockRuntime` — reopening the same path from a second runtime to + /// exercise resume, or inspecting the database directly. + /// + /// The caller owns the path, and therefore owns cleaning it up. The + /// default leaves nothing to clean up. + pub fn with_store(mut self, store: SqliteRuntimeStore) -> Self { + self.store = Some(store); + self + } + + /// Replaces the runtime policy used by the scripted runtime. + pub fn with_policy(mut self, policy: RuntimePolicy) -> Self { + self.policy = policy; + self + } + + /// Installs a tool authorizer, so a scripted run can exercise the + /// permission flow. + /// + /// Without one the session authorizer allows every call unconditionally + /// and [`SessionEvent::PermissionRequested`](crate::SessionEvent) is never + /// emitted — which makes "does this host ask before it writes?" impossible + /// to test against a mock. + /// Installs a pre-execution hook, so a scripted run can exercise the + /// interception path. + /// + /// The sibling of [`with_tool_authorizer`](Self::with_tool_authorizer): + /// without one, nothing ever consults a hook, so a host can test that its + /// own hook logic is correct but not that the runtime actually calls it. + pub fn with_pre_hook(mut self, hook: impl PreExecutionHook + 'static) -> Self { + self.pre_hook = Some(Box::new(hook)); + self + } + + pub fn with_tool_authorizer(mut self, authorizer: impl ToolAuthorizer + 'static) -> Self { + self.tool_authorizer = Some(Box::new(authorizer)); + self + } + + pub fn push_turn(mut self, turn: MockTurn) -> Self { + self.turns.push(turn); + self + } + + pub fn text(self, text: impl Into) -> Self { + self.push_turn(MockTurn::Text(text.into())) + } + + pub fn stream_text(self, chunks: I) -> Self + where + I: IntoIterator, + S: Into, + { + self.push_turn(MockTurn::StreamText( + chunks.into_iter().map(Into::into).collect(), + )) + } + + pub fn tool_calls(self, calls: I) -> Self + where + I: IntoIterator, + { + self.push_turn(MockTurn::ToolCalls(calls.into_iter().collect())) + } + + pub fn failure(self, error: ProviderError) -> Self { + self.push_turn(MockTurn::Failure(error)) + } + + pub fn build(self) -> Result { + let provider = ScriptedProvider::new(self.model.provider.clone(), vec![self.model.clone()]); + provider.push_turns(self.turns); + + let mut builder = Runtime::builder() + .with_runtime_identifier(self.runtime_identifier) + .with_policy(self.policy) + .with_provider_instance(provider.clone()); + + // A scripted runtime is ephemeral by definition, so its store is too. + // The default used to be a SQLite file named after the current + // nanosecond in the system temp directory, which left one file behind + // per mock and — when two mocks read the same tick — handed both the + // same database, where the second one's agent lease was already held. + builder = match self.store { + Some(store) => builder.with_store(store), + None => builder.with_store(VolatileRuntimeStore::new()), + }; + + if let Some(hook) = self.pre_hook { + builder = builder.with_pre_hook(hook); + } + + if let Some(authorizer) = self.tool_authorizer { + builder = builder.with_tool_authorizer(authorizer); + } + + let runtime = builder.build()?; + + Ok(MockRuntime { + runtime, + provider, + model: self.model, + }) + } +} + +#[derive(Clone)] +struct ScriptedProvider { + kind: ProviderId, + models: Vec, + turns: Arc>>, + requests: Arc>>>, +} + +impl ScriptedProvider { + fn new(kind: ProviderId, models: Vec) -> Self { + Self { + kind, + models, + turns: Arc::new(Mutex::new(VecDeque::new())), + requests: Arc::new(Mutex::new(Vec::new())), + } + } + + fn push_turns(&self, turns: Vec) { + let mut queue = self.turns.lock().expect("mock turn queue poisoned"); + queue.extend(turns); + } + + fn recorded_requests(&self) -> Vec> { + self.requests + .lock() + .expect("mock request log poisoned") + .clone() + } +} + +#[async_trait] +impl Provider for ScriptedProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.kind.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(self.models.clone()) + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.requests + .lock() + .expect("mock request log poisoned") + .push(request.into_owned()); + let turn = self + .turns + .lock() + .expect("mock turn queue poisoned") + .pop_front() + .unwrap_or_else(|| panic!("no scripted turn remaining for mock runtime")); + + match turn { + MockTurn::Text(text) => Ok(response_stream( + &self.models[0], + Response { + id: format!("mock-response-{}", now_nanos()), + model: self.models[0].id.clone(), + role: Role::Assistant, + content: vec![ContentBlock::text(text)], + stop_reason: None, + usage: None, + }, + )), + MockTurn::StreamText(chunks) => Ok(streaming_text_response(&self.models[0], chunks)), + MockTurn::ToolCalls(calls) => Ok(response_stream( + &self.models[0], + Response { + id: format!("mock-response-{}", now_nanos()), + model: self.models[0].id.clone(), + role: Role::Assistant, + content: calls + .into_iter() + .enumerate() + .map(|(index, call)| ContentBlock::ToolUse { + id: call.id.unwrap_or_else(|| format!("tool-{}", index + 1)), + name: call.name, + input: call.input, + }) + .collect(), + stop_reason: Some("tool_use".to_string()), + usage: None, + }, + )), + MockTurn::Failure(error) => Err(error), + } + } +} + +fn response_stream(model: &ModelInfo, response: Response) -> ProviderEventStream { + let _ = model; + provider_event_stream_from_response(response) +} + +fn streaming_text_response(model: &ModelInfo, chunks: Vec) -> ProviderEventStream { + let (tx, rx) = mpsc::unbounded_channel(); + + tx.send(Ok(ProviderEvent::MessageStarted { + id: format!("mock-response-{}", now_nanos()), + model: model.id.clone(), + role: Role::Assistant, + })) + .expect("mock runtime message start receiver dropped"); + tx.send(Ok(ProviderEvent::ContentBlockStarted { + index: 0, + kind: crate::provider::ContentBlockStart::Text, + })) + .expect("mock runtime content start receiver dropped"); + + for chunk in chunks { + tx.send(Ok(ProviderEvent::ContentBlockDelta { + index: 0, + delta: crate::provider::ContentBlockDelta::Text(chunk), + })) + .expect("mock runtime content delta receiver dropped"); + } + + tx.send(Ok(ProviderEvent::ContentBlockStopped { index: 0 })) + .expect("mock runtime content stop receiver dropped"); + tx.send(Ok(ProviderEvent::MessageDelta { + stop_reason: None, + usage: None, + })) + .expect("mock runtime message delta receiver dropped"); + tx.send(Ok(ProviderEvent::MessageStopped)) + .expect("mock runtime message stop receiver dropped"); + + rx +} + +fn now_nanos() -> u128 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + use crate::{ + Agent, + agent::{AgentConfig, ToolProfile}, + provider::Message, + tool::{ParallelToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSpec}, + }; + + struct EchoTool; + + #[async_trait] + impl ToolDefinition for EchoTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("echo_tool") + .description("Echo a canned result") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .build() + } + } + + #[async_trait] + impl ToolExecutor for EchoTool { + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + Ok("echoed".to_string()) + } + } + + async fn spawn_agent(mock: &MockRuntime) -> Agent { + mock.runtime() + .spawn("mock-agent", mock.model()) + .expect("spawn mock agent") + } + + #[tokio::test] + async fn mock_runtime_replays_text_turns() { + let mock = MockRuntime::builder() + .text("hello from mock") + .build() + .unwrap(); + let mut agent = spawn_agent(&mock).await; + + let message = agent.send(vec![ContentBlock::text("hi")]).await.unwrap(); + + assert_eq!( + message, + Message::assistant(ContentBlock::text("hello from mock")) + ); + } + + #[tokio::test] + async fn mock_runtime_replays_streaming_text_turns() { + let mock = MockRuntime::builder() + .stream_text(["hello", " ", "world"]) + .build() + .unwrap(); + let mut agent = spawn_agent(&mock).await; + + let message = agent.send(vec![ContentBlock::text("hi")]).await.unwrap(); + + assert_eq!(message.text(), "hello world"); + } + + #[tokio::test] + async fn mock_runtime_surfaces_provider_failures() { + let mock = MockRuntime::builder() + .failure(ProviderError::InvalidResponse("boom".to_string())) + .build() + .unwrap(); + let mut agent = spawn_agent(&mock).await; + + let error = agent + .send(vec![ContentBlock::text("hi")]) + .await + .unwrap_err(); + + assert!(matches!(error, RuntimeError::FailedToStreamResponse(_))); + } + + #[tokio::test] + async fn mock_runtime_can_script_tool_call_turns() { + let mock = MockRuntime::builder() + .tool_calls([MockToolCall::new("echo_tool", json!({}))]) + .text("done") + .build() + .unwrap(); + mock.runtime().register_tool(EchoTool); + let mut agent = spawn_agent(&mock).await; + + let message = agent + .send(vec![ContentBlock::text("run the tool")]) + .await + .unwrap(); + + assert_eq!(message.text(), "done"); + assert_eq!(mock.recorded_requests().await.len(), 2); + } + + #[tokio::test] + async fn mock_runtime_supports_runtime_assembly_assertions() { + let mock = MockRuntime::builder().text("done").build().unwrap(); + mock.runtime().register_tool(EchoTool); + let mut agent = mock + .runtime() + .spawn_with_config( + "mock-agent", + mock.model(), + AgentConfig { + tool_profile: ToolProfile::only(["echo_tool"]), + ..Default::default() + }, + ) + .expect("spawn mock agent"); + + let message = agent.send(vec![ContentBlock::text("hi")]).await.unwrap(); + + assert_eq!(message.text(), "done"); + + let requests = mock.recorded_requests().await; + let tool_names = requests[0] + .tools + .iter() + .map(|tool| tool.name.as_str()) + .collect::>(); + assert_eq!(tool_names, vec!["echo_tool"]); + } + + /// An authorizer that prompts for everything, so the session authorizer + /// has something to raise. + struct AlwaysPrompts; + + #[async_trait] + impl crate::tool::ToolAuthorizer for AlwaysPrompts { + async fn authorize( + &self, + _request: &crate::tool::ToolAuthorizationRequest, + ) -> Result { + Ok(crate::tool::ToolAuthorizationDecision::prompt("ask first")) + } + } + + #[tokio::test] + async fn a_mock_runtime_can_exercise_the_permission_path() { + let mock = MockRuntime::builder() + .with_tool_authorizer(AlwaysPrompts) + // A builtin, so the call really reaches the tool layer and + // therefore the authorizer. + .tool_calls(vec![MockToolCall::new( + "files", + json!({"operations": [{"op": "list", "path": "."}]}), + )]) + .text("done") + .build() + .unwrap(); + + let mut session = mock.runtime().create_session("test", mock.model()).unwrap(); + + let mut events = session.subscribe(); + let permissions = session.permission_handle(); + let asked = Arc::new(Mutex::new(false)); + let saw = Arc::clone(&asked); + + // Answer whatever is asked, so the turn is not left blocked forever. + let watcher = tokio::spawn(async move { + while let Ok(event) = events.recv().await { + if let crate::SessionEvent::PermissionRequested { request_id, .. } = event { + *saw.lock().unwrap() = true; + let _ = permissions.resolve_permission( + &request_id, + crate::session::PermissionDecision::deny(), + ); + } + } + }); + + let _ = session.append_turn(vec![ContentBlock::text("go")]).await; + watcher.abort(); + + assert!( + *asked.lock().unwrap(), + "an authorizer installed on the mock must reach the session's permission flow" + ); + } + + /// Denies one named tool, so a scripted run can prove the runtime really + /// consults the hook rather than that the hook's own logic is correct. + struct DenyTool(&'static str); + + #[async_trait] + impl crate::runtime::PreExecutionHook for DenyTool { + async fn pre_tool_execution( + &self, + context: &crate::runtime::PreExecutionContext, + ) -> Result { + if context.tool_name == self.0 { + Ok(crate::runtime::HookDecision::Deny( + "not this one".to_string(), + )) + } else { + Ok(crate::runtime::HookDecision::Allow) + } + } + } + + #[tokio::test] + async fn a_mock_runtime_can_exercise_the_interception_path() { + let mock = MockRuntime::builder() + .with_pre_hook(DenyTool("files")) + .tool_calls(vec![MockToolCall::new( + "files", + json!({"operations": [{"op": "list", "path": "."}]}), + )]) + .text("done") + .build() + .unwrap(); + + let mut session = mock.runtime().create_session("test", mock.model()).unwrap(); + + let _ = session.append_turn(vec![ContentBlock::text("go")]).await; + + let blocked = session.replay().items().iter().any(|item| { + item.message.as_ref().is_some_and(|message| { + message.content.iter().any(|block| { + matches!(block, ContentBlock::ToolResult { content, .. } + if content.to_string().contains("not this one")) + }) + }) + }); + + assert!( + blocked, + "a hook installed on the mock must actually be consulted by the runtime" + ); + } + + /// How many mock-runtime databases the system temp directory holds right + /// now. Counted as a delta rather than asserted at zero, because a machine + /// that ran the old default has thousands of them left over and the point + /// is that this run adds none. + fn mock_runtime_files_in_temp() -> usize { + let Ok(entries) = std::fs::read_dir(std::env::temp_dir()) else { + return 0; + }; + entries + .filter_map(Result::ok) + .filter(|entry| { + entry + .file_name() + .to_string_lossy() + .starts_with("mentra-mock-runtime") + }) + .count() + } + + /// The default store used to be a SQLite file in the system temp + /// directory, one per mock, named after the current nanosecond and never + /// deleted. A full downstream suite left dozens behind per run; one dev + /// machine had accumulated 38,782. + #[tokio::test] + async fn a_default_mock_runtime_writes_nothing_to_disk() { + let before = mock_runtime_files_in_temp(); + + let mock = MockRuntime::builder().text("hello").build().unwrap(); + let mut agent = spawn_agent(&mock).await; + agent + .send(vec![ContentBlock::text("hi")]) + .await + .expect("a scripted turn completes without a store on disk"); + + assert_eq!( + mock_runtime_files_in_temp(), + before, + "a mock runtime must leave nothing in {}", + std::env::temp_dir().display() + ); + } + + /// Two mocks built inside one nanosecond used to be handed the same + /// database file. Agent ids are unique only within a process, so two test + /// binaries running concurrently could mint the same id against that + /// shared file — and the second `spawn` failed with `LeaseUnavailable`, + /// the mechanism behind a downstream flake nobody could reproduce. + /// Independent stores make the collision unreachable rather than rare. + #[tokio::test] + async fn mock_runtimes_built_back_to_back_do_not_share_a_store() { + let first = MockRuntime::builder() + .runtime_identifier("shared-identifier") + .text("from the first") + .build() + .unwrap(); + let second = MockRuntime::builder() + .runtime_identifier("shared-identifier") + .text("from the second") + .build() + .unwrap(); + + // Both spawns take an agent lease. Against one shared store, the + // second is the one that would be refused. + let mut first_agent = spawn_agent(&first).await; + let mut second_agent = spawn_agent(&second).await; + + assert_eq!( + first_agent + .send(vec![ContentBlock::text("hi")]) + .await + .expect("the first mock runs") + .text(), + "from the first" + ); + assert_eq!( + second_agent + .send(vec![ContentBlock::text("hi")]) + .await + .expect("the second mock runs") + .text(), + "from the second" + ); + + // Same runtime identifier, so anything they shared would show up here. + for mock in [&first, &second] { + let agents = mock + .runtime() + .list_persisted_agents("shared-identifier") + .expect("lists persisted agents"); + assert_eq!( + agents.len(), + 1, + "each mock keeps its own store, so it sees only its own agent" + ); + } + } +} diff --git a/vendor/mentra/src/tool.rs b/vendor/mentra/src/tool.rs new file mode 100644 index 0000000..fa1422a --- /dev/null +++ b/vendor/mentra/src/tool.rs @@ -0,0 +1,211 @@ +mod authorization; +/// Bash command validation — safety checks before shell execution. +pub mod bash_validation; +mod builtin; +mod coding; +mod context; +mod descriptor; +mod files; +pub(crate) mod internal; +mod model; +mod orchestrator; +pub(crate) mod paging; +mod runtime; +mod truncation; + +use std::{collections::HashMap, sync::Arc}; + +pub use authorization::{ + ToolAuthorizationDecision, ToolAuthorizationOutcome, ToolAuthorizationPreview, + ToolAuthorizationRequest, ToolAuthorizer, +}; +pub use descriptor::{ + ProviderToolSpec, RuntimeToolDescriptor, RuntimeToolDescriptorBuilder, ToolApprovalCategory, + ToolCapability, ToolDurability, ToolExecutionCategory, ToolExecutionMode, ToolLoadingPolicy, + ToolSideEffectLevel, +}; +pub use mentra_provider::ToolResultContent; +pub use model::{ + ExecutableTool, ParallelToolContext, ToolCall, ToolContext, ToolDefinition, ToolExecutor, + ToolOutput, ToolResult, ToolSpec, +}; +pub(crate) use runtime::ToolRuntime; + +pub(crate) use builtin::ReadToolResultTool; +use builtin::{BackgroundRunTool, CheckBackgroundTool, LoadSkillTool, ShellTool}; +use coding::{EditTool, GlobTool, GrepTool, ListTool, ReadTool, WriteTool}; +use files::FilesTool; + +/// Selects which builtin file-tool surface a runtime exposes. +/// +/// [`Batched`](Self::Batched) preserves the historical `files` tool exactly. +/// [`Split`](Self::Split) exposes model-conventional `read`, `ls`, `grep`, +/// `glob`, `write`, and `edit` tools. [`Both`](Self::Both) exposes both +/// surfaces over the same workspace engine. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub enum FileToolProfile { + #[default] + Batched, + Split, + Both, +} + +#[derive(Clone)] +struct RegisteredTool { + descriptor: RuntimeToolDescriptor, + handler: Arc, +} + +#[derive(Clone, Default)] +/// Registry of tools available to a runtime instance. +pub struct ToolRegistry { + tools: HashMap, + provider_specs: Arc<[ProviderToolSpec]>, +} + +impl ToolRegistry { + /// Registers a tool implementation and refreshes the cached tool specs. + pub fn register_tool(&mut self, tool: T) + where + T: ExecutableTool + 'static, + { + let handler: Arc = Arc::new(tool); + let descriptor = handler.descriptor(); + self.tools.insert( + descriptor.provider.name.clone(), + RegisteredTool { + descriptor, + handler, + }, + ); + self.refresh_provider_specs(); + } + + /// Returns the provider-facing tool specifications. + pub fn tools(&self) -> Arc<[ProviderToolSpec]> { + Arc::clone(&self.provider_specs) + } + + /// Returns a tool handler by name. + pub fn get_tool(&self, name: &str) -> Option> { + self.tools.get(name).map(|tool| Arc::clone(&tool.handler)) + } + + pub fn get_tool_descriptor(&self, name: &str) -> Option { + self.tools.get(name).map(|tool| tool.descriptor.clone()) + } + + pub(crate) fn unregister_tool(&mut self, name: &str) -> bool { + let removed = self.tools.remove(name).is_some(); + if removed { + self.refresh_provider_specs(); + } + removed + } + + fn refresh_provider_specs(&mut self) { + self.provider_specs = self + .tools + .values() + .map(|tool| tool.descriptor.provider.clone()) + .collect::>() + .into(); + } +} + +impl ToolRegistry { + pub(crate) fn register_skill_tool(&mut self) { + self.register_tool(LoadSkillTool); + } + + pub(crate) fn register_builtin_tools(&mut self, file_tools: FileToolProfile) { + self.register_tool(ShellTool); + self.register_tool(BackgroundRunTool); + self.register_tool(CheckBackgroundTool); + self.configure_file_tools(file_tools); + } + + pub(crate) fn configure_file_tools(&mut self, profile: FileToolProfile) { + for name in ["files", "read", "ls", "grep", "glob", "write", "edit"] { + self.tools.remove(name); + } + + if matches!(profile, FileToolProfile::Batched | FileToolProfile::Both) { + self.register_tool(FilesTool); + } + if matches!(profile, FileToolProfile::Split | FileToolProfile::Both) { + self.register_tool(ReadTool); + self.register_tool(ListTool); + self.register_tool(GrepTool); + self.register_tool(GlobTool); + self.register_tool(WriteTool); + self.register_tool(EditTool); + } + self.refresh_provider_specs(); + } +} + +#[cfg(test)] +mod tests { + use std::{borrow::Cow, collections::BTreeMap}; + + use serde_json::json; + + use super::*; + + #[test] + fn builtin_shell_and_files_tools_serialize_as_non_strict_responses_functions() { + let mut registry = ToolRegistry::default(); + registry.register_builtin_tools(FileToolProfile::default()); + + let request = mentra_provider::Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(Vec::new()), + tools: Cow::Owned(registry.tools().to_vec()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: mentra_provider::ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value( + mentra_provider::responses::model::ResponsesRequest::try_from(request) + .expect("built-in tools should serialize for Responses"), + ) + .expect("responses request should serialize"); + let tools = payload["tools"] + .as_array() + .expect("tools should be a json array"); + + for name in ["shell", "background_run", "files"] { + let tool = tools + .iter() + .find(|tool| tool["name"] == json!(name)) + .unwrap_or_else(|| panic!("{name} tool should be serialized")); + assert_eq!(tool["type"], "function"); + assert_eq!(tool["strict"], false); + } + } + + #[test] + fn file_tool_profiles_replace_the_eager_builtin_surface() { + let mut registry = ToolRegistry::default(); + registry.register_builtin_tools(FileToolProfile::Batched); + assert!(registry.get_tool("files").is_some()); + assert!(registry.get_tool("read").is_none()); + + registry.configure_file_tools(FileToolProfile::Split); + assert!(registry.get_tool("files").is_none()); + for name in ["read", "ls", "grep", "glob", "write", "edit"] { + assert!(registry.get_tool(name).is_some(), "missing {name}"); + } + + registry.configure_file_tools(FileToolProfile::Both); + assert!(registry.get_tool("files").is_some()); + for name in ["read", "ls", "grep", "glob", "write", "edit"] { + assert!(registry.get_tool(name).is_some(), "missing {name}"); + } + } +} diff --git a/vendor/mentra/src/tool/authorization.rs b/vendor/mentra/src/tool/authorization.rs new file mode 100644 index 0000000..53d0163 --- /dev/null +++ b/vendor/mentra/src/tool/authorization.rs @@ -0,0 +1,120 @@ +use std::path::PathBuf; +use std::time::Duration; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +use crate::{ + runtime::RuntimeError, + tool::{ + ToolApprovalCategory, ToolCapability, ToolDurability, ToolExecutionCategory, + ToolSideEffectLevel, + }, +}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ToolAuthorizationOutcome { + Allow, + Prompt, + Deny, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ToolAuthorizationPreview { + pub working_directory: PathBuf, + pub capabilities: Vec, + pub side_effect_level: ToolSideEffectLevel, + pub durability: ToolDurability, + pub execution_category: ToolExecutionCategory, + pub approval_category: ToolApprovalCategory, + pub raw_input: Value, + pub structured_input: Value, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ToolAuthorizationDecision { + pub outcome: ToolAuthorizationOutcome, + pub reason: Option, +} + +impl ToolAuthorizationDecision { + pub fn allow() -> Self { + Self { + outcome: ToolAuthorizationOutcome::Allow, + reason: None, + } + } + + pub fn prompt(reason: impl Into) -> Self { + Self { + outcome: ToolAuthorizationOutcome::Prompt, + reason: Some(reason.into()), + } + } + + pub fn deny(reason: impl Into) -> Self { + Self { + outcome: ToolAuthorizationOutcome::Deny, + reason: Some(reason.into()), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ToolAuthorizationRequest { + pub agent_id: String, + pub agent_name: String, + pub model: String, + pub history_len: usize, + pub tool_call_id: String, + pub tool_name: String, + pub preview: ToolAuthorizationPreview, +} + +#[async_trait] +pub trait ToolAuthorizer: Send + Sync { + async fn authorize( + &self, + request: &ToolAuthorizationRequest, + ) -> Result; + + fn timeout(&self) -> Option { + None + } +} + +/// Forwards to the authorizer inside. +/// +/// Lets a caller hold an authorizer it chose at runtime — one of several, or +/// none — and still hand it to anything taking `impl ToolAuthorizer`, without +/// each caller writing this impl itself. +#[async_trait] +impl ToolAuthorizer for Box { + async fn authorize( + &self, + request: &ToolAuthorizationRequest, + ) -> Result { + (**self).authorize(request).await + } + + fn timeout(&self) -> Option { + (**self).timeout() + } +} + +/// Forwards to the authorizer inside, for a shared one. +#[async_trait] +impl ToolAuthorizer for std::sync::Arc { + async fn authorize( + &self, + request: &ToolAuthorizationRequest, + ) -> Result { + (**self).authorize(request).await + } + + fn timeout(&self) -> Option { + (**self).timeout() + } +} diff --git a/vendor/mentra/src/tool/bash_validation.rs b/vendor/mentra/src/tool/bash_validation.rs new file mode 100644 index 0000000..c4b0d4e --- /dev/null +++ b/vendor/mentra/src/tool/bash_validation.rs @@ -0,0 +1,852 @@ +//! Bash command validation — safety checks before shell execution. +//! +//! Provides heuristic classification and validation of shell commands to detect +//! destructive, write, or suspicious operations before they execute. + +use std::path::Path; + +use crate::tool::ToolAuthorizationOutcome; + +/// Result of validating a bash command before execution. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ValidationResult { + /// Command is safe to execute. + Allow, + /// Command should be blocked with the given reason. + Block { reason: String }, + /// Command requires user confirmation with the given warning. + Warn { message: String }, +} + +impl ValidationResult { + pub(crate) fn authorization_outcome(&self) -> ToolAuthorizationOutcome { + match self { + Self::Allow => ToolAuthorizationOutcome::Allow, + Self::Block { .. } => ToolAuthorizationOutcome::Deny, + Self::Warn { .. } => ToolAuthorizationOutcome::Prompt, + } + } + + pub(crate) fn reason(&self) -> Option<&str> { + match self { + Self::Allow => None, + Self::Block { reason } => Some(reason), + Self::Warn { message } => Some(message), + } + } +} + +/// Semantic classification of a bash command's intent. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CommandIntent { + ReadOnly, + Write, + Destructive, + Network, + ProcessManagement, + PackageManagement, + SystemAdmin, + Unknown, +} + +// --------------------------------------------------------------------------- +// Command lists +// --------------------------------------------------------------------------- + +const WRITE_COMMANDS: &[&str] = &[ + "cp", "mv", "rm", "mkdir", "rmdir", "touch", "chmod", "chown", "chgrp", "ln", "install", "tee", + "truncate", "shred", "mkfifo", "mknod", "dd", +]; + +const STATE_MODIFYING_COMMANDS: &[&str] = &[ + "apt", + "apt-get", + "yum", + "dnf", + "pacman", + "brew", + "pip", + "pip3", + "npm", + "yarn", + "pnpm", + "bun", + "cargo", + "gem", + "go", + "rustup", + "docker", + "systemctl", + "service", + "mount", + "umount", + "kill", + "pkill", + "killall", + "reboot", + "shutdown", + "halt", + "poweroff", + "useradd", + "userdel", + "usermod", + "groupadd", + "groupdel", + "crontab", + "at", +]; + +const WRITE_REDIRECTIONS: &[&str] = &[">", ">>", ">&"]; + +const READ_ONLY_COMMANDS: &[&str] = &[ + "ls", + "cat", + "head", + "tail", + "less", + "more", + "wc", + "sort", + "uniq", + "grep", + "egrep", + "fgrep", + "find", + "which", + "whereis", + "whatis", + "man", + "file", + "stat", + "du", + "df", + "free", + "uptime", + "uname", + "hostname", + "whoami", + "id", + "groups", + "env", + "printenv", + "echo", + "printf", + "date", + "cal", + "bc", + "expr", + "test", + "true", + "false", + "pwd", + "tree", + "diff", + "cmp", + "md5sum", + "sha256sum", + "sha1sum", + "xxd", + "od", + "hexdump", + "strings", + "readlink", + "realpath", + "basename", + "dirname", + "seq", + "tput", + "column", + "jq", + "yq", + "xargs", + "tr", + "cut", + "paste", + "awk", + "sed", + "rg", +]; + +const NETWORK_COMMANDS: &[&str] = &[ + "curl", + "wget", + "ssh", + "scp", + "rsync", + "ftp", + "sftp", + "nc", + "ncat", + "telnet", + "ping", + "traceroute", + "dig", + "nslookup", + "host", + "whois", + "ifconfig", + "ip", + "netstat", + "ss", + "nmap", +]; + +const PROCESS_COMMANDS: &[&str] = &[ + "kill", "pkill", "killall", "ps", "top", "htop", "bg", "fg", "jobs", "nohup", "disown", "wait", + "nice", "renice", +]; + +const PACKAGE_COMMANDS: &[&str] = &[ + "apt", "apt-get", "yum", "dnf", "pacman", "brew", "pip", "pip3", "npm", "yarn", "pnpm", "bun", + "cargo", "gem", "go", "rustup", "snap", "flatpak", +]; + +const SYSTEM_ADMIN_COMMANDS: &[&str] = &[ + "sudo", + "su", + "chroot", + "mount", + "umount", + "fdisk", + "parted", + "systemctl", + "service", + "iptables", + "ufw", + "sysctl", + "crontab", + "at", + "useradd", + "userdel", + "usermod", + "groupadd", + "groupdel", + "passwd", + "visudo", +]; + +const ALWAYS_DESTRUCTIVE_COMMANDS: &[&str] = &["shred", "wipefs"]; + +const DESTRUCTIVE_PATTERNS: &[(&str, &str)] = &[ + ( + "rm -rf /", + "Recursive forced deletion at root — this will destroy the system", + ), + ("rm -rf ~", "Recursive forced deletion of home directory"), + ( + "rm -rf *", + "Recursive forced deletion of all files in current directory", + ), + ("rm -rf .", "Recursive forced deletion of current directory"), + ( + "mkfs", + "Filesystem creation will destroy existing data on the device", + ), + ( + "dd if=", + "Direct disk write — can overwrite partitions or devices", + ), + ("> /dev/sd", "Writing to raw disk device"), + ( + "chmod -R 777", + "Recursively setting world-writable permissions", + ), + ("chmod -R 000", "Recursively removing all permissions"), + (":(){ :|:& };:", "Fork bomb — will crash the system"), +]; + +const GIT_READ_ONLY_SUBCOMMANDS: &[&str] = &[ + "status", + "log", + "diff", + "show", + "branch", + "tag", + "stash", + "remote", + "fetch", + "ls-files", + "ls-tree", + "cat-file", + "rev-parse", + "describe", + "shortlog", + "blame", + "bisect", + "reflog", + "config", +]; + +const SYSTEM_PATHS: &[&str] = &[ + "/etc/", "/usr/", "/var/", "/boot/", "/sys/", "/proc/", "/dev/", "/sbin/", "/lib/", "/opt/", +]; + +// --------------------------------------------------------------------------- +// Public API +// --------------------------------------------------------------------------- + +/// Check if a command is destructive and should be warned about. +#[must_use] +pub fn check_destructive(command: &str) -> ValidationResult { + for &(pattern, warning) in DESTRUCTIVE_PATTERNS { + if command.contains(pattern) { + return ValidationResult::Warn { + message: format!("Destructive command detected: {warning}"), + }; + } + } + + let first = extract_first_command(command); + for &cmd in ALWAYS_DESTRUCTIVE_COMMANDS { + if first == cmd { + return ValidationResult::Warn { + message: format!( + "Command '{cmd}' is inherently destructive and may cause data loss" + ), + }; + } + } + + if command.contains("rm ") && command.contains("-r") && command.contains("-f") { + return ValidationResult::Warn { + message: "Recursive forced deletion detected — verify the target path is correct" + .to_string(), + }; + } + + ValidationResult::Allow +} + +/// Check if a command targets paths outside the workspace. +#[must_use] +pub fn check_workspace_escape(command: &str) -> ValidationResult { + let first = extract_first_command(command); + let is_write_cmd = WRITE_COMMANDS.contains(&first.as_str()) + || STATE_MODIFYING_COMMANDS.contains(&first.as_str()); + + if !is_write_cmd { + return ValidationResult::Allow; + } + + for sys_path in SYSTEM_PATHS { + if command.contains(sys_path) { + return ValidationResult::Warn { + message: + "Command appears to target files outside the workspace — requires elevated permission" + .to_string(), + }; + } + } + + ValidationResult::Allow +} + +/// Validate path patterns in a command. +#[must_use] +pub fn validate_paths(command: &str, workspace: &Path) -> ValidationResult { + if command.contains("../") { + let workspace_str = workspace.to_string_lossy(); + if !command.contains(&*workspace_str) { + return ValidationResult::Warn { + message: "Command contains directory traversal pattern '../' — verify the target path resolves within the workspace".to_string(), + }; + } + } + + if command.contains("~/") || command.contains("$HOME") { + return ValidationResult::Warn { + message: + "Command references home directory — verify it stays within the workspace scope" + .to_string(), + }; + } + + ValidationResult::Allow +} + +/// Validate sed-specific safety. +#[must_use] +pub fn validate_sed(command: &str, read_only: bool) -> ValidationResult { + let first = extract_first_command(command); + if first != "sed" { + return ValidationResult::Allow; + } + + if read_only && command.contains(" -i") { + return ValidationResult::Block { + reason: "sed -i (in-place editing) is not allowed in read-only mode".to_string(), + }; + } + + ValidationResult::Allow +} + +/// Validate a command for read-only mode. +#[must_use] +pub fn validate_read_only(command: &str) -> ValidationResult { + let first_command = extract_first_command(command); + + for &write_cmd in WRITE_COMMANDS { + if first_command == write_cmd { + return ValidationResult::Block { + reason: format!( + "Command '{write_cmd}' modifies the filesystem and is not allowed in read-only mode" + ), + }; + } + } + + for &state_cmd in STATE_MODIFYING_COMMANDS { + if first_command == state_cmd { + return ValidationResult::Block { + reason: format!( + "Command '{state_cmd}' modifies system state and is not allowed in read-only mode" + ), + }; + } + } + + if first_command == "sudo" { + let inner = extract_sudo_inner(command); + if !inner.is_empty() { + let inner_result = validate_read_only(inner); + if inner_result != ValidationResult::Allow { + return inner_result; + } + } + } + + for &redir in WRITE_REDIRECTIONS { + if command.contains(redir) { + return ValidationResult::Block { + reason: format!( + "Command contains write redirection '{redir}' which is not allowed in read-only mode" + ), + }; + } + } + + if first_command == "git" { + return validate_git_read_only(command); + } + + ValidationResult::Allow +} + +/// Classify the semantic intent of a bash command. +#[must_use] +pub fn classify_command(command: &str) -> CommandIntent { + let first = extract_first_command(command); + + if READ_ONLY_COMMANDS.contains(&first.as_str()) { + if first == "sed" && command.contains(" -i") { + return CommandIntent::Write; + } + return CommandIntent::ReadOnly; + } + + if ALWAYS_DESTRUCTIVE_COMMANDS.contains(&first.as_str()) || first == "rm" { + return CommandIntent::Destructive; + } + + if WRITE_COMMANDS.contains(&first.as_str()) { + return CommandIntent::Write; + } + + if NETWORK_COMMANDS.contains(&first.as_str()) { + return CommandIntent::Network; + } + + if PROCESS_COMMANDS.contains(&first.as_str()) { + return CommandIntent::ProcessManagement; + } + + if PACKAGE_COMMANDS.contains(&first.as_str()) { + return CommandIntent::PackageManagement; + } + + if SYSTEM_ADMIN_COMMANDS.contains(&first.as_str()) { + return CommandIntent::SystemAdmin; + } + + if first == "git" { + return classify_git_command(command); + } + + CommandIntent::Unknown +} + +/// Run the full validation pipeline on a bash command. +#[must_use] +pub fn validate_command(command: &str, workspace: &Path, read_only: bool) -> ValidationResult { + if read_only { + let result = validate_read_only(command); + if result != ValidationResult::Allow { + return result; + } + } + + let result = validate_sed(command, read_only); + if result != ValidationResult::Allow { + return result; + } + + let result = check_destructive(command); + if result != ValidationResult::Allow { + return result; + } + + let result = check_workspace_escape(command); + if result != ValidationResult::Allow { + return result; + } + + validate_paths(command, workspace) +} + +// --------------------------------------------------------------------------- +// Internal helpers +// --------------------------------------------------------------------------- + +fn validate_git_read_only(command: &str) -> ValidationResult { + let parts: Vec<&str> = command.split_whitespace().collect(); + let subcommand = parts.iter().skip(1).find(|p| !p.starts_with('-')); + + match subcommand { + Some(&sub) if GIT_READ_ONLY_SUBCOMMANDS.contains(&sub) => ValidationResult::Allow, + Some(&sub) => ValidationResult::Block { + reason: format!( + "Git subcommand '{sub}' modifies repository state and is not allowed in read-only mode" + ), + }, + None => ValidationResult::Allow, + } +} + +fn classify_git_command(command: &str) -> CommandIntent { + let parts: Vec<&str> = command.split_whitespace().collect(); + let subcommand = parts.iter().skip(1).find(|p| !p.starts_with('-')); + match subcommand { + Some(&sub) if GIT_READ_ONLY_SUBCOMMANDS.contains(&sub) => CommandIntent::ReadOnly, + _ => CommandIntent::Write, + } +} + +fn extract_first_command(command: &str) -> String { + let trimmed = command.trim(); + let mut remaining = trimmed; + + // Skip leading environment variable assignments. + loop { + let next = remaining.trim_start(); + if let Some(eq_pos) = next.find('=') { + let before_eq = &next[..eq_pos]; + if !before_eq.is_empty() + && before_eq + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '_') + { + let after_eq = &next[eq_pos + 1..]; + if let Some(space) = find_end_of_value(after_eq) { + remaining = &after_eq[space..]; + continue; + } + return String::new(); + } + } + break; + } + + remaining + .split_whitespace() + .next() + .unwrap_or("") + .to_string() +} + +fn extract_sudo_inner(command: &str) -> &str { + let parts: Vec<&str> = command.split_whitespace().collect(); + let sudo_idx = parts.iter().position(|&p| p == "sudo"); + match sudo_idx { + Some(idx) => { + let rest = &parts[idx + 1..]; + for &part in rest { + if !part.starts_with('-') { + let offset = command.find(part).unwrap_or(0); + return &command[offset..]; + } + } + "" + } + None => "", + } +} + +fn find_end_of_value(s: &str) -> Option { + let s = s.trim_start(); + if s.is_empty() { + return None; + } + + let first = s.as_bytes()[0]; + if first == b'"' || first == b'\'' { + let quote = first; + let mut i = 1; + while i < s.len() { + if s.as_bytes()[i] == quote && (i == 0 || s.as_bytes()[i - 1] != b'\\') { + i += 1; + while i < s.len() && !s.as_bytes()[i].is_ascii_whitespace() { + i += 1; + } + return if i < s.len() { Some(i) } else { None }; + } + i += 1; + } + None + } else { + s.find(char::is_whitespace) + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + use std::path::PathBuf; + + #[test] + fn classify_read_only_commands() { + assert_eq!(classify_command("ls -la"), CommandIntent::ReadOnly); + assert_eq!(classify_command("cat file.txt"), CommandIntent::ReadOnly); + assert_eq!( + classify_command("grep -r pattern ."), + CommandIntent::ReadOnly + ); + assert_eq!( + classify_command("find . -name '*.rs'"), + CommandIntent::ReadOnly + ); + assert_eq!(classify_command("rg pattern"), CommandIntent::ReadOnly); + } + + #[test] + fn classify_write_commands() { + assert_eq!(classify_command("cp a.txt b.txt"), CommandIntent::Write); + assert_eq!(classify_command("mv old.txt new.txt"), CommandIntent::Write); + assert_eq!(classify_command("mkdir -p /tmp/dir"), CommandIntent::Write); + } + + #[test] + fn classify_destructive_commands() { + assert_eq!( + classify_command("rm -rf /tmp/x"), + CommandIntent::Destructive + ); + assert_eq!( + classify_command("shred /dev/sda"), + CommandIntent::Destructive + ); + } + + #[test] + fn classify_network_commands() { + assert_eq!( + classify_command("curl https://example.com"), + CommandIntent::Network + ); + assert_eq!(classify_command("wget file.zip"), CommandIntent::Network); + } + + #[test] + fn classify_sed_inplace_as_write() { + assert_eq!( + classify_command("sed -i 's/old/new/' file.txt"), + CommandIntent::Write + ); + } + + #[test] + fn classify_sed_stdout_as_read_only() { + assert_eq!( + classify_command("sed 's/old/new/' file.txt"), + CommandIntent::ReadOnly + ); + } + + #[test] + fn classify_git_status_as_read_only() { + assert_eq!(classify_command("git status"), CommandIntent::ReadOnly); + assert_eq!( + classify_command("git log --oneline"), + CommandIntent::ReadOnly + ); + } + + #[test] + fn classify_git_push_as_write() { + assert_eq!( + classify_command("git push origin main"), + CommandIntent::Write + ); + } + + #[test] + fn blocks_rm_in_read_only() { + assert!(matches!( + validate_read_only("rm -rf /tmp/x"), + ValidationResult::Block { reason } if reason.contains("rm") + )); + } + + #[test] + fn allows_ls_in_read_only() { + assert_eq!(validate_read_only("ls -la"), ValidationResult::Allow); + } + + #[test] + fn blocks_write_redirect_in_read_only() { + assert!(matches!( + validate_read_only("echo hello > file.txt"), + ValidationResult::Block { reason } if reason.contains("redirection") + )); + } + + #[test] + fn blocks_sudo_rm_in_read_only() { + assert!(matches!( + validate_read_only("sudo rm -rf /tmp/x"), + ValidationResult::Block { reason } if reason.contains("rm") + )); + } + + #[test] + fn blocks_git_push_in_read_only() { + assert!(matches!( + validate_read_only("git push origin main"), + ValidationResult::Block { reason } if reason.contains("push") + )); + } + + #[test] + fn allows_git_status_in_read_only() { + assert_eq!(validate_read_only("git status"), ValidationResult::Allow); + } + + #[test] + fn warns_rm_rf_root() { + assert!(matches!( + check_destructive("rm -rf /"), + ValidationResult::Warn { message } if message.contains("root") + )); + } + + #[test] + fn warns_fork_bomb() { + assert!(matches!( + check_destructive(":(){ :|:& };:"), + ValidationResult::Warn { message } if message.contains("Fork bomb") + )); + } + + #[test] + fn allows_safe_destructive_check() { + assert_eq!(check_destructive("ls -la"), ValidationResult::Allow); + } + + #[test] + fn warns_system_paths() { + assert!(matches!( + check_workspace_escape("cp file.txt /etc/config"), + ValidationResult::Warn { .. } + )); + } + + #[test] + fn allows_local_write() { + assert_eq!( + check_workspace_escape("cp file.txt ./backup/"), + ValidationResult::Allow + ); + } + + #[test] + fn warns_directory_traversal() { + let workspace = PathBuf::from("/workspace/project"); + assert!(matches!( + validate_paths("cat ../../../etc/passwd", &workspace), + ValidationResult::Warn { message } if message.contains("traversal") + )); + } + + #[test] + fn warns_home_reference() { + let workspace = PathBuf::from("/workspace"); + assert!(matches!( + validate_paths("cat ~/.ssh/id_rsa", &workspace), + ValidationResult::Warn { message } if message.contains("home directory") + )); + } + + #[test] + fn full_pipeline_blocks_write_in_read_only() { + let workspace = PathBuf::from("/workspace"); + assert!(matches!( + validate_command("rm -rf /tmp/x", &workspace, true), + ValidationResult::Block { .. } + )); + } + + #[test] + fn full_pipeline_warns_destructive() { + let workspace = PathBuf::from("/workspace"); + assert!(matches!( + validate_command("rm -rf /", &workspace, false), + ValidationResult::Warn { .. } + )); + } + + #[test] + fn full_pipeline_allows_safe_read() { + let workspace = PathBuf::from("/workspace"); + assert_eq!( + validate_command("ls -la", &workspace, true), + ValidationResult::Allow + ); + } + + #[test] + fn validation_results_map_to_authorization_outcomes() { + assert_eq!( + ValidationResult::Allow.authorization_outcome(), + ToolAuthorizationOutcome::Allow + ); + assert_eq!( + ValidationResult::Warn { + message: "review".to_string(), + } + .authorization_outcome(), + ToolAuthorizationOutcome::Prompt + ); + assert_eq!( + ValidationResult::Block { + reason: "blocked".to_string(), + } + .authorization_outcome(), + ToolAuthorizationOutcome::Deny + ); + } + + #[test] + fn extracts_command_from_env_prefix() { + assert_eq!(extract_first_command("FOO=bar ls -la"), "ls"); + assert_eq!(extract_first_command("A=1 B=2 echo hello"), "echo"); + } + + #[test] + fn extracts_plain_command() { + assert_eq!(extract_first_command("grep -r pattern ."), "grep"); + } +} diff --git a/vendor/mentra/src/tool/builtin.rs b/vendor/mentra/src/tool/builtin.rs new file mode 100644 index 0000000..8a8f11a --- /dev/null +++ b/vendor/mentra/src/tool/builtin.rs @@ -0,0 +1,10 @@ +#[path = "builtin/read_only.rs"] +mod read_only; +#[path = "builtin/read_tool_result.rs"] +mod read_tool_result; +#[path = "builtin/shell.rs"] +mod shell; + +pub use read_only::{CheckBackgroundTool, LoadSkillTool}; +pub(crate) use read_tool_result::ReadToolResultTool; +pub use shell::{BackgroundRunTool, ShellTool}; diff --git a/vendor/mentra/src/tool/builtin/read_only.rs b/vendor/mentra/src/tool/builtin/read_only.rs new file mode 100644 index 0000000..4879ec4 --- /dev/null +++ b/vendor/mentra/src/tool/builtin/read_only.rs @@ -0,0 +1,83 @@ +use async_trait::async_trait; +use serde_json::json; + +use crate::tool::{ + ParallelToolContext, RuntimeToolDescriptor, ToolApprovalCategory, ToolCapability, + ToolDefinition, ToolDurability, ToolExecutionCategory, ToolExecutor, ToolResult, + ToolSideEffectLevel, +}; + +pub struct CheckBackgroundTool; +pub struct LoadSkillTool; + +fn check_background_descriptor() -> RuntimeToolDescriptor { + RuntimeToolDescriptor::builder("check_background") + .description("Check one background task by ID, or list all background tasks when omitted.") + .input_schema(json!({ + "type": "object", + "properties": { + "task_id": { + "type": "string", + "description": "Optional background task ID to inspect" + } + } + })) + .capability(ToolCapability::ReadOnly) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .approval_category(ToolApprovalCategory::ReadOnly) + .build() +} + +fn load_skill_descriptor() -> RuntimeToolDescriptor { + RuntimeToolDescriptor::builder("load_skill") + .description("Load the full body of a named skill when it is relevant.") + .input_schema(json!({ + "type": "object", + "properties": { + "name": { + "type": "string", + "description": "Name of the skill to load" + } + }, + "required": ["name"] + })) + .capabilities([ToolCapability::SkillLoad, ToolCapability::ReadOnly]) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .approval_category(ToolApprovalCategory::ReadOnly) + .build() +} + +impl ToolDefinition for CheckBackgroundTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + check_background_descriptor() + } +} + +#[async_trait] +impl ToolExecutor for CheckBackgroundTool { + async fn execute(&self, ctx: ParallelToolContext, input: serde_json::Value) -> ToolResult { + let task_id = input.get("task_id").and_then(|value| value.as_str()); + ctx.check_background_task(task_id) + } +} + +impl ToolDefinition for LoadSkillTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + load_skill_descriptor() + } +} + +#[async_trait] +impl ToolExecutor for LoadSkillTool { + async fn execute(&self, ctx: ParallelToolContext, input: serde_json::Value) -> ToolResult { + let name = input + .get("name") + .and_then(|value| value.as_str()) + .ok_or_else(|| "Skill name is required".to_string())?; + ctx.load_skill(name) + } +} diff --git a/vendor/mentra/src/tool/builtin/read_tool_result.rs b/vendor/mentra/src/tool/builtin/read_tool_result.rs new file mode 100644 index 0000000..ab85476 --- /dev/null +++ b/vendor/mentra/src/tool/builtin/read_tool_result.rs @@ -0,0 +1,97 @@ +//! The `read_tool_result` built-in: reads further windows of a tool result +//! that was too large to deliver whole. +//! +//! It serves only the results this agent's own run retained, it never +//! re-executes the tool that produced them (so an expensive or +//! side-effectful call is not repeated to read more of its output), and its +//! own output is bounded by the agent's `page_bytes`, so reading can never +//! reintroduce the overflow paging exists to prevent. + +use async_trait::async_trait; +use serde::Deserialize; +use serde_json::{Value, json}; + +use crate::tool::{ + ToolApprovalCategory, ToolCapability, ToolContext, ToolDefinition, ToolDurability, + ToolExecutionCategory, ToolExecutor, ToolResult, ToolSideEffectLevel, ToolSpec, + paging::{READ_TOOL_RESULT_TOOL, ToolResultPager}, +}; + +pub(crate) struct ReadToolResultTool; + +#[derive(Deserialize)] +struct ReadToolResultInput { + tool_use_id: String, + start_line: usize, +} + +impl ToolDefinition for ReadToolResultTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder(READ_TOOL_RESULT_TOOL) + .description( + "Read the next window of a tool result that was too large to deliver whole. \ + Pass the tool_use_id and start_line printed in that result's paging trailer. \ + Line numbers are absolute over the full result, so they mean the same thing \ + in every window. This reads retained output only — it never re-runs the tool \ + that produced it.", + ) + .input_schema(json!({ + "type": "object", + "properties": { + "tool_use_id": { + "type": "string", + "description": "The tool_use_id printed in the paging trailer." + }, + "start_line": { + "type": "integer", + "minimum": 1, + "description": "1-based absolute line to start the window at." + } + }, + "required": ["tool_use_id", "start_line"] + })) + .capability(ToolCapability::ReadOnly) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + // Exclusive rather than ReadOnlyParallel despite reading nothing + // but memory: the retained results are agent state, and only the + // exclusive lane's `ToolContext` carries the agent. + .execution_category(ToolExecutionCategory::ExclusiveLocalMutation) + .approval_category(ToolApprovalCategory::ReadOnly) + .build() + } +} + +#[async_trait] +impl ToolExecutor for ReadToolResultTool { + async fn execute_mut(&self, ctx: ToolContext<'_>, input: Value) -> ToolResult { + let request: ReadToolResultInput = serde_json::from_value(input).map_err(|error| { + format!( + "{READ_TOOL_RESULT_TOOL} expects a tool_use_id string and a start_line integer \ + of 1 or greater: {error}" + ) + })?; + if request.start_line == 0 { + return Err(format!( + "{READ_TOOL_RESULT_TOOL} start_line is 1-based; 0 is not a line" + )); + } + + let Some(paging) = ctx.tool_result_paging() else { + return Err(format!( + "{READ_TOOL_RESULT_TOOL} is unavailable: tool-result paging is not enabled \ + for this agent" + )); + }; + let Some(full) = ctx.paged_tool_result(&request.tool_use_id) else { + return Err(format!( + "no retained result for tool_use_id \"{}\". Only results large enough to be \ + paged are retained, and only for this agent's current run — use the \ + tool_use_id printed in a paging trailer.", + request.tool_use_id + )); + }; + + Ok(ToolResultPager::new(paging).window(&request.tool_use_id, &full, request.start_line)) + } +} diff --git a/vendor/mentra/src/tool/builtin/shell.rs b/vendor/mentra/src/tool/builtin/shell.rs new file mode 100644 index 0000000..e0ffe5f --- /dev/null +++ b/vendor/mentra/src/tool/builtin/shell.rs @@ -0,0 +1,278 @@ +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::tool::{ + ParallelToolContext, RuntimeToolDescriptor, ToolApprovalCategory, ToolAuthorizationPreview, + ToolCapability, ToolDefinition, ToolDurability, ToolExecutionCategory, ToolExecutor, + ToolResult, ToolSideEffectLevel, context::RuntimeContext, +}; + +pub struct ShellTool; +pub struct BackgroundRunTool; + +struct ShellCommandInput<'a> { + command: String, + working_directory: Option<&'a str>, + justification: Option, + requested_timeout: Option, +} + +fn parse_shell_command_input<'a>(input: &'a Value) -> Result, String> { + let command = input + .get("command") + .and_then(|value| value.as_str()) + .ok_or_else(|| "Command is required".to_string())? + .to_string(); + let working_directory = input + .get("workingDirectory") + .and_then(|value| value.as_str()); + let justification = input + .get("justification") + .and_then(|value| value.as_str()) + .map(ToOwned::to_owned); + let requested_timeout = input + .get("timeoutMs") + .and_then(|value| value.as_u64()) + .map(std::time::Duration::from_millis); + + Ok(ShellCommandInput { + command, + working_directory, + justification, + requested_timeout, + }) +} + +fn shell_input_schema(include_timeout: bool) -> Value { + let mut properties = serde_json::Map::from_iter([ + ( + "command".to_string(), + json!({ + "type": "string", + "description": "Shell command to execute" + }), + ), + ( + "workingDirectory".to_string(), + json!({ + "type": "string", + "description": "Optional directory to run inside" + }), + ), + ( + "justification".to_string(), + json!({ + "type": "string", + "description": "Optional explanation surfaced when approval is required" + }), + ), + ]); + if include_timeout { + properties.insert( + "timeoutMs".to_string(), + json!({ + "type": "integer", + "description": "Optional timeout override in milliseconds" + }), + ); + } + Value::Object(serde_json::Map::from_iter([ + ("type".to_string(), json!("object")), + ("properties".to_string(), Value::Object(properties)), + ("required".to_string(), json!(["command"])), + ])) +} + +fn shell_descriptor(background: bool) -> RuntimeToolDescriptor { + let (name, description, capabilities, durability, execution_category, approval_category) = + if background { + ( + "background_run", + "Start a shell command in the background and return a task ID immediately.", + vec![ + ToolCapability::BackgroundExec, + ToolCapability::FilesystemWrite, + ], + ToolDurability::Persistent, + ToolExecutionCategory::BackgroundJob, + ToolApprovalCategory::Background, + ) + } else { + ( + "shell", + "Execute a single local shell command.", + vec![ToolCapability::ProcessExec, ToolCapability::FilesystemWrite], + ToolDurability::Ephemeral, + ToolExecutionCategory::ExclusiveLocalMutation, + ToolApprovalCategory::Process, + ) + }; + + RuntimeToolDescriptor::builder(name) + .description(description) + .input_schema(shell_input_schema(!background)) + .capabilities(capabilities) + .side_effect_level(ToolSideEffectLevel::Process) + .durability(durability) + .execution_category(execution_category) + .approval_category(approval_category) + .build() +} + +fn shell_authorization_preview( + ctx: &ParallelToolContext, + input: &Value, + background: bool, + descriptor: RuntimeToolDescriptor, +) -> Result { + let ShellCommandInput { + command, + working_directory, + justification, + requested_timeout, + } = parse_shell_command_input(input)?; + let working_directory = ctx.resolve_working_directory(working_directory)?; + let validation = ctx.shell_validation(&command)?; + + Ok(ToolAuthorizationPreview { + working_directory: working_directory.clone(), + capabilities: descriptor.capabilities, + side_effect_level: descriptor.side_effect_level, + durability: descriptor.durability, + execution_category: descriptor.execution_category, + approval_category: descriptor.approval_category, + raw_input: input.clone(), + structured_input: json!({ + "kind": if background { "background_run" } else { "shell" }, + "command": command, + "working_directory": working_directory, + "timeout_ms": requested_timeout.map(|timeout| timeout.as_millis()), + "justification": justification, + "background": background, + "validation": { + "mode": validation.mode.as_str(), + "intent": validation.intent_name(), + "outcome": validation.outcome, + "reason": validation.reason(), + }, + }), + }) +} + +fn emit_output_progress(ctx: &C, output: &crate::runtime::CommandOutput) { + if !output.stdout.is_empty() { + for line in output.stdout.lines() { + ctx.emit_progress(format!("stdout: {line}")); + } + } + if !output.stderr.is_empty() { + for line in output.stderr.lines() { + ctx.emit_progress(format!("stderr: {line}")); + } + } +} + +async fn execute_shell_command(ctx: &C, input: Value) -> ToolResult +where + C: RuntimeContext + Sync, +{ + let ShellCommandInput { + command, + working_directory, + justification, + requested_timeout, + } = parse_shell_command_input(&input)?; + let working_directory = ctx.resolve_working_directory(working_directory)?; + let output = ctx + .execute_shell_command(command, justification, requested_timeout, working_directory) + .await?; + + emit_output_progress(ctx, &output); + + if output.success() { + if !output.stdout.is_empty() { + Ok(output.stdout) + } else { + Ok(output.stderr) + } + } else { + let message = if !output.stderr.trim().is_empty() { + output.stderr + } else if !output.stdout.trim().is_empty() { + output.stdout + } else if output.timed_out { + "Command timed out after the configured limit".to_string() + } else { + format!( + "Command exited with status {}", + output + .status_code + .map(|code| code.to_string()) + .unwrap_or_else(|| "unknown".to_string()) + ) + }; + Err(message) + } +} + +async fn execute_background_command(ctx: &C, input: Value) -> ToolResult +where + C: RuntimeContext + Sync, +{ + let ShellCommandInput { + command, + working_directory, + justification, + .. + } = parse_shell_command_input(&input)?; + let working_directory = ctx.resolve_working_directory(working_directory)?; + let task = ctx.start_background_task(command, justification, None, working_directory)?; + Ok(format!( + "Started background task {} in {} for `{}`", + task.id, + task.cwd.display(), + task.command + )) +} + +impl ToolDefinition for ShellTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + shell_descriptor(false) + } +} + +#[async_trait] +impl ToolExecutor for ShellTool { + fn authorization_preview( + &self, + ctx: &ParallelToolContext, + input: &Value, + ) -> Result { + shell_authorization_preview(ctx, input, false, self.descriptor()) + } + + async fn execute_mut(&self, ctx: crate::tool::ToolContext<'_>, input: Value) -> ToolResult { + execute_shell_command(&ctx, input).await + } +} + +impl ToolDefinition for BackgroundRunTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + shell_descriptor(true) + } +} + +#[async_trait] +impl ToolExecutor for BackgroundRunTool { + fn authorization_preview( + &self, + ctx: &ParallelToolContext, + input: &Value, + ) -> Result { + shell_authorization_preview(ctx, input, true, self.descriptor()) + } + + async fn execute_mut(&self, ctx: crate::tool::ToolContext<'_>, input: Value) -> ToolResult { + execute_background_command(&ctx, input).await + } +} diff --git a/vendor/mentra/src/tool/coding.rs b/vendor/mentra/src/tool/coding.rs new file mode 100644 index 0000000..61538cf --- /dev/null +++ b/vendor/mentra/src/tool/coding.rs @@ -0,0 +1,264 @@ +#[path = "coding/execution.rs"] +mod execution; +#[path = "coding/input.rs"] +mod input; + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::tool::{ + ParallelToolContext, RuntimeToolDescriptor, ToolApprovalCategory, ToolCapability, ToolContext, + ToolDefinition, ToolDurability, ToolExecutionCategory, ToolExecutor, ToolOutput, ToolResult, + ToolSideEffectLevel, +}; + +use execution::{ + execute_edit, execute_glob, execute_grep, execute_list, execute_read, execute_write, +}; + +pub(crate) struct ReadTool; +pub(crate) struct ListTool; +pub(crate) struct GrepTool; +pub(crate) struct GlobTool; +pub(crate) struct WriteTool; +pub(crate) struct EditTool; + +fn read_only_descriptor( + name: &str, + description: &str, + input_schema: Value, +) -> RuntimeToolDescriptor { + RuntimeToolDescriptor::builder(name) + .description(description) + .input_schema(input_schema) + .capabilities([ToolCapability::ReadOnly, ToolCapability::FilesystemRead]) + .side_effect_level(ToolSideEffectLevel::None) + .durability(ToolDurability::ReplaySafe) + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .approval_category(ToolApprovalCategory::ReadOnly) + .build() +} + +fn mutation_descriptor( + name: &str, + description: &str, + input_schema: Value, +) -> RuntimeToolDescriptor { + RuntimeToolDescriptor::builder(name) + .description(description) + .input_schema(input_schema) + .capability(ToolCapability::FilesystemWrite) + .side_effect_level(ToolSideEffectLevel::LocalState) + .durability(ToolDurability::Ephemeral) + .execution_category(ToolExecutionCategory::ExclusiveLocalMutation) + .approval_category(ToolApprovalCategory::Filesystem) + .build() +} + +impl ToolDefinition for ReadTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + read_only_descriptor( + "read", + "Read a UTF-8 text file from the workspace with line numbers.", + json!({ + "type": "object", + "properties": { + "path": { "type": "string" }, + "file_path": { "type": "string" }, + "offset": { "type": "integer", "minimum": 1 }, + "limit": { "type": "integer", "minimum": 0 } + }, + "anyOf": [ + { "required": ["path"] }, + { "required": ["file_path"] } + ] + }), + ) + } +} + +impl ToolDefinition for ListTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + read_only_descriptor( + "ls", + "List files and directories within the workspace.", + json!({ + "type": "object", + "properties": { + "path": { "type": "string" }, + "depth": { "type": "integer", "minimum": 0 }, + "limit": { "type": "integer", "minimum": 0 } + } + }), + ) + } +} + +impl ToolDefinition for GrepTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + read_only_descriptor( + "grep", + "Search workspace text files with a regular expression or literal pattern.", + json!({ + "type": "object", + "properties": { + "pattern": { "type": "string" }, + "path": { "type": "string" }, + "glob": { "type": "string" }, + "ignore_case": { "type": "boolean" }, + "literal": { "type": "boolean" }, + "context": { "type": "integer", "minimum": 0 }, + "multiline": { "type": "boolean" }, + "limit": { "type": "integer", "minimum": 0 } + }, + "required": ["pattern"] + }), + ) + } +} + +impl ToolDefinition for GlobTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + read_only_descriptor( + "glob", + "Find workspace files whose relative paths match a glob pattern.", + json!({ + "type": "object", + "properties": { + "pattern": { "type": "string" }, + "path": { "type": "string" }, + "limit": { "type": "integer", "minimum": 0 } + }, + "required": ["pattern"] + }), + ) + } +} + +impl ToolDefinition for WriteTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + mutation_descriptor( + "write", + "Create or overwrite a UTF-8 text file within the workspace.", + json!({ + "type": "object", + "properties": { + "path": { "type": "string" }, + "file_path": { "type": "string" }, + "content": { "type": "string" } + }, + "required": ["content"], + "anyOf": [ + { "required": ["path"] }, + { "required": ["file_path"] } + ] + }), + ) + } +} + +impl ToolDefinition for EditTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + mutation_descriptor( + "edit", + "Replace one or more uniquely matched text blocks in a workspace file.", + json!({ + "type": "object", + "properties": { + "path": { "type": "string" }, + "file_path": { "type": "string" }, + "edits": { + "type": "array", + "items": { + "type": "object", + "properties": { + "old_string": { "type": "string" }, + "new_string": { "type": "string" } + }, + "required": ["old_string", "new_string"] + }, + "minItems": 1 + }, + "replace_all": { "type": "boolean" } + }, + "required": ["edits"], + "anyOf": [ + { "required": ["path"] }, + { "required": ["file_path"] } + ] + }), + ) + } +} + +#[async_trait] +impl ToolExecutor for ReadTool { + async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult { + execute_read(ctx, input).await + } +} + +#[async_trait] +impl ToolExecutor for ListTool { + async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult { + execute_list(ctx, input).await + } +} + +#[async_trait] +impl ToolExecutor for GrepTool { + async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult { + execute_grep(ctx, input).await + } +} + +#[async_trait] +impl ToolExecutor for GlobTool { + async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult { + execute_glob(ctx, input).await + } +} + +#[async_trait] +impl ToolExecutor for WriteTool { + async fn execute_mut(&self, ctx: ToolContext<'_>, input: Value) -> ToolResult { + execute_write(ctx.into(), input).await + } +} + +#[async_trait] +impl ToolExecutor for EditTool { + async fn execute_mut_output( + &self, + ctx: ToolContext<'_>, + input: Value, + ) -> Result { + execute_edit(ctx.into(), input).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn split_tools_have_static_scheduler_categories() { + for descriptor in [ + ReadTool.descriptor(), + ListTool.descriptor(), + GrepTool.descriptor(), + GlobTool.descriptor(), + ] { + assert_eq!( + descriptor.execution_category, + ToolExecutionCategory::ReadOnlyParallel + ); + } + for descriptor in [WriteTool.descriptor(), EditTool.descriptor()] { + assert_eq!( + descriptor.execution_category, + ToolExecutionCategory::ExclusiveLocalMutation + ); + } + } +} diff --git a/vendor/mentra/src/tool/coding/execution.rs b/vendor/mentra/src/tool/coding/execution.rs new file mode 100644 index 0000000..18d8cc1 --- /dev/null +++ b/vendor/mentra/src/tool/coding/execution.rs @@ -0,0 +1,118 @@ +use serde_json::{Value, json}; + +use crate::tool::{ParallelToolContext, ToolOutput, ToolResult}; + +use super::input::{parse_edit, parse_glob, parse_grep, parse_list, parse_read, parse_write}; +use crate::tool::files::workspace::{EditOutcome, TextEdit, WorkspaceEditor}; + +pub(super) async fn execute_read(ctx: ParallelToolContext, input: Value) -> ToolResult { + let input = parse_read(input)?; + with_editor(ctx, move |editor| { + editor.read( + input.path, + input.offset.unwrap_or(1), + input.limit.unwrap_or(2_000), + ) + }) + .await +} + +pub(super) async fn execute_list(ctx: ParallelToolContext, input: Value) -> ToolResult { + let input = parse_list(input)?; + with_editor(ctx, move |editor| { + editor.list( + input.path.unwrap_or_else(|| ".".to_string()), + input.depth.unwrap_or(1), + input.limit.unwrap_or(200), + ) + }) + .await +} + +pub(super) async fn execute_grep(ctx: ParallelToolContext, input: Value) -> ToolResult { + let input = parse_grep(input)?; + let options = input.search_options(); + with_editor(ctx, move |editor| { + editor.grep( + input.path.unwrap_or_else(|| ".".to_string()), + &input.pattern, + options, + input.limit.unwrap_or(200), + ) + }) + .await +} + +pub(super) async fn execute_glob(ctx: ParallelToolContext, input: Value) -> ToolResult { + let input = parse_glob(input)?; + with_editor(ctx, move |editor| { + editor.glob( + input.path.unwrap_or_else(|| ".".to_string()), + &input.pattern, + input.limit.unwrap_or(200), + ) + }) + .await +} + +pub(super) async fn execute_write(ctx: ParallelToolContext, input: Value) -> ToolResult { + let input = parse_write(input)?; + with_editor(ctx, move |editor| editor.write(input.path, input.content)).await +} + +pub(super) async fn execute_edit( + ctx: ParallelToolContext, + input: Value, +) -> Result { + let input = parse_edit(input)?; + let edits = input + .edits + .into_iter() + .map(TextEdit::from) + .collect::>(); + let outcome = with_editor(ctx, move |editor| { + editor.edit(input.path, edits, input.replace_all) + }) + .await?; + Ok(edit_output(outcome)) +} + +fn edit_output(outcome: EditOutcome) -> ToolOutput { + let block_label = if outcome.replacement_count == 1 { + "block" + } else { + "blocks" + }; + ToolOutput::text(format!( + "Replaced {} {block_label} in {}", + outcome.replacement_count, outcome.display_path + )) + .with_details(json!({ + "diff": outcome.diff, + "patch": outcome.patch, + "first_changed_line": outcome.first_changed_line, + })) +} + +async fn with_editor(ctx: ParallelToolContext, operation: F) -> Result +where + T: Send + 'static, + F: FnOnce(&mut WorkspaceEditor) -> Result + Send + 'static, +{ + let base_dir = ctx.runtime.agent_config(&ctx.agent_id)?.base_dir; + let working_directory = ctx + .runtime + .resolve_working_directory(&ctx.agent_id, None) + .unwrap_or_else(|_| ctx.working_directory().to_path_buf()); + let agent_id = ctx.agent_id; + let runtime = ctx.runtime; + + tokio::task::spawn_blocking(move || { + let mut editor = WorkspaceEditor::new(agent_id, runtime, base_dir, working_directory); + let result = operation(&mut editor)?; + editor.commit()?; + Ok(result) + }) + .await + .map_err(|error| format!("Coding tool task failed: {error}"))? +} diff --git a/vendor/mentra/src/tool/coding/input.rs b/vendor/mentra/src/tool/coding/input.rs new file mode 100644 index 0000000..b1b11f1 --- /dev/null +++ b/vendor/mentra/src/tool/coding/input.rs @@ -0,0 +1,202 @@ +use serde::Deserialize; +use serde_json::{Map, Value}; + +use crate::tool::files::workspace::{SearchOptions, TextEdit}; + +#[derive(Debug, Deserialize)] +pub(super) struct ReadInput { + #[serde(alias = "file_path", alias = "filePath")] + pub(super) path: String, + pub(super) offset: Option, + pub(super) limit: Option, +} + +#[derive(Debug, Deserialize)] +pub(super) struct ListInput { + #[serde(default)] + pub(super) path: Option, + pub(super) depth: Option, + pub(super) limit: Option, +} + +#[derive(Debug, Deserialize)] +pub(super) struct GrepInput { + pub(super) pattern: String, + #[serde(default)] + pub(super) path: Option, + #[serde(default)] + pub(super) glob: Option, + #[serde(default, alias = "ignoreCase")] + pub(super) ignore_case: bool, + #[serde(default)] + pub(super) literal: bool, + #[serde(default)] + pub(super) context: usize, + #[serde(default)] + pub(super) multiline: bool, + pub(super) limit: Option, +} + +impl GrepInput { + pub(super) fn search_options(&self) -> SearchOptions { + SearchOptions { + file_glob: self.glob.clone(), + ignore_case: self.ignore_case, + literal: self.literal, + context: self.context, + multiline: self.multiline, + max_line_chars: Some(500), + } + } +} + +#[derive(Debug, Deserialize)] +pub(super) struct GlobInput { + pub(super) pattern: String, + #[serde(default)] + pub(super) path: Option, + pub(super) limit: Option, +} + +#[derive(Debug, Deserialize)] +pub(super) struct WriteInput { + #[serde(alias = "file_path", alias = "filePath")] + pub(super) path: String, + pub(super) content: String, +} + +#[derive(Debug, Deserialize)] +pub(super) struct EditInput { + #[serde(alias = "file_path", alias = "filePath")] + pub(super) path: String, + pub(super) edits: Vec, + #[serde(default, alias = "replaceAll")] + pub(super) replace_all: bool, +} + +#[derive(Debug, Deserialize)] +pub(super) struct EditSpec { + #[serde(alias = "oldText", alias = "old")] + old_string: String, + #[serde(alias = "newText", alias = "new")] + new_string: String, +} + +impl From for TextEdit { + fn from(value: EditSpec) -> Self { + Self { + old_string: value.old_string, + new_string: value.new_string, + } + } +} + +pub(super) fn parse_read(input: Value) -> Result { + parse(input, "read") +} + +pub(super) fn parse_list(input: Value) -> Result { + parse(input, "ls") +} + +pub(super) fn parse_grep(input: Value) -> Result { + parse(input, "grep") +} + +pub(super) fn parse_glob(input: Value) -> Result { + parse(input, "glob") +} + +pub(super) fn parse_write(input: Value) -> Result { + parse(input, "write") +} + +pub(super) fn parse_edit(input: Value) -> Result { + let mut object = input + .as_object() + .cloned() + .ok_or_else(|| "Invalid edit input: expected an object".to_string())?; + normalize_edits(&mut object)?; + parse(Value::Object(object), "edit") +} + +fn normalize_edits(object: &mut Map) -> Result<(), String> { + if let Some(Value::String(encoded)) = object.get("edits") { + let decoded: Value = serde_json::from_str(encoded) + .map_err(|error| format!("Invalid edit input: edits JSON string: {error}"))?; + object.insert("edits".to_string(), normalize_edit_collection(decoded)?); + } else if let Some(edits) = object.get("edits").cloned() { + object.insert("edits".to_string(), normalize_edit_collection(edits)?); + } else { + let old = take_first(object, &["old_string", "oldText", "old"]); + let new = take_first(object, &["new_string", "newText", "new"]); + if old.is_some() || new.is_some() { + let mut edit = Map::new(); + if let Some(old) = old { + edit.insert("old_string".to_string(), old); + } + if let Some(new) = new { + edit.insert("new_string".to_string(), new); + } + object.insert("edits".to_string(), Value::Array(vec![Value::Object(edit)])); + } + } + Ok(()) +} + +fn normalize_edit_collection(value: Value) -> Result { + match value { + Value::Array(_) => Ok(value), + Value::Object(_) => Ok(Value::Array(vec![value])), + _ => Err("Invalid edit input: edits must be an array, object, or JSON string".to_string()), + } +} + +fn take_first(object: &mut Map, keys: &[&str]) -> Option { + keys.iter().find_map(|key| object.remove(*key)) +} + +fn parse(input: Value, tool: &str) -> Result +where + T: for<'de> Deserialize<'de>, +{ + serde_json::from_value(input).map_err(|error| format!("Invalid {tool} input: {error}")) +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + + #[test] + fn edit_accepts_json_encoded_edits_and_camel_case_aliases() { + let parsed = parse_edit(json!({ + "filePath": "src/lib.rs", + "edits": r#"[{"oldText":"before","newText":"after"}]"#, + "replaceAll": true + })) + .expect("parse edit"); + + assert_eq!(parsed.path, "src/lib.rs"); + assert!(parsed.replace_all); + assert_eq!(parsed.edits.len(), 1); + assert_eq!(parsed.edits[0].old_string, "before"); + assert_eq!(parsed.edits[0].new_string, "after"); + } + + #[test] + fn edit_accepts_legacy_top_level_single_edit() { + let parsed = parse_edit(json!({ + "file_path": "src/lib.rs", + "old_string": "before", + "new_string": "after" + })) + .expect("parse edit"); + + assert_eq!(parsed.path, "src/lib.rs"); + assert_eq!(parsed.edits.len(), 1); + assert_eq!(parsed.edits[0].old_string, "before"); + assert_eq!(parsed.edits[0].new_string, "after"); + } +} diff --git a/vendor/mentra/src/tool/context.rs b/vendor/mentra/src/tool/context.rs new file mode 100644 index 0000000..a7c3dd8 --- /dev/null +++ b/vendor/mentra/src/tool/context.rs @@ -0,0 +1,96 @@ +use async_trait::async_trait; + +use crate::tool::{ParallelToolContext, ToolContext}; + +#[async_trait] +pub(crate) trait RuntimeContext { + fn resolve_working_directory( + &self, + working_directory: Option<&str>, + ) -> Result; + async fn execute_shell_command( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: std::path::PathBuf, + ) -> Result; + fn start_background_task( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: std::path::PathBuf, + ) -> Result; + fn emit_progress(&self, progress: String); +} + +#[async_trait] +impl RuntimeContext for ToolContext<'_> { + fn resolve_working_directory( + &self, + working_directory: Option<&str>, + ) -> Result { + self.resolve_working_directory(working_directory) + } + + async fn execute_shell_command( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: std::path::PathBuf, + ) -> Result { + self.execute_shell_command(command, justification, requested_timeout, cwd) + .await + } + + fn start_background_task( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: std::path::PathBuf, + ) -> Result { + self.start_background_task(command, justification, requested_timeout, cwd) + } + + fn emit_progress(&self, progress: String) { + self.emit_progress(progress); + } +} + +#[async_trait] +impl RuntimeContext for ParallelToolContext { + fn resolve_working_directory( + &self, + working_directory: Option<&str>, + ) -> Result { + self.resolve_working_directory(working_directory) + } + + async fn execute_shell_command( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: std::path::PathBuf, + ) -> Result { + self.execute_shell_command(command, justification, requested_timeout, cwd) + .await + } + + fn start_background_task( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: std::path::PathBuf, + ) -> Result { + self.start_background_task(command, justification, requested_timeout, cwd) + } + + fn emit_progress(&self, progress: String) { + self.emit_progress(progress); + } +} diff --git a/vendor/mentra/src/tool/descriptor.rs b/vendor/mentra/src/tool/descriptor.rs new file mode 100644 index 0000000..df3a004 --- /dev/null +++ b/vendor/mentra/src/tool/descriptor.rs @@ -0,0 +1,284 @@ +use std::{ops::Deref, time::Duration}; + +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +pub use mentra_provider::{ + ProviderToolKind, ToolLoadingPolicy, ToolSpec as ProviderToolSpec, + ToolSpecBuilder as ProviderToolSpecBuilder, +}; + +/// High-level capability labels used for runtime metadata and policy decisions. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum ToolCapability { + ReadOnly, + FilesystemRead, + FilesystemWrite, + ProcessExec, + BackgroundExec, + TaskMutation, + TeamCoordination, + Delegation, + ContextCompaction, + SkillLoad, + Custom(String), +} + +/// Declares how much side effect a tool may have when executed. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +pub enum ToolSideEffectLevel { + #[default] + None, + LocalState, + Process, + External, +} + +/// Declares whether a tool call is safe to replay or persist. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +pub enum ToolDurability { + #[default] + Ephemeral, + Persistent, + ReplaySafe, +} + +/// Declares which scheduler/orchestrator lane a tool call belongs to. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +pub enum ToolExecutionCategory { + ReadOnlyParallel, + #[default] + ExclusiveLocalMutation, + ExclusivePersistentMutation, + BackgroundJob, + Delegation, +} + +impl ToolExecutionCategory { + pub fn allows_parallel(self) -> bool { + matches!(self, Self::ReadOnlyParallel) + } +} + +/// Backward-compatible parallel/exclusive view of execution semantics. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +pub enum ToolExecutionMode { + #[default] + Exclusive, + Parallel, +} + +impl From for ToolExecutionMode { + fn from(value: ToolExecutionCategory) -> Self { + if value.allows_parallel() { + Self::Parallel + } else { + Self::Exclusive + } + } +} + +/// Coarse authorization grouping for runtime policy and review systems. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +pub enum ToolApprovalCategory { + #[default] + Default, + ReadOnly, + Filesystem, + Process, + Background, + Delegation, +} + +/// Runtime-facing descriptor that wraps the provider-visible tool definition. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct RuntimeToolDescriptor { + pub provider: ProviderToolSpec, + pub capabilities: Vec, + pub side_effect_level: ToolSideEffectLevel, + pub durability: ToolDurability, + pub execution_category: ToolExecutionCategory, + pub approval_category: ToolApprovalCategory, + pub execution_timeout: Option, + /// Marks this tool as a terminal action (see + /// [`RuntimeToolDescriptorBuilder::terminal`]). A terminal-marked tool is + /// never scheduled in a parallel batch: the scheduler coerces it to an + /// exclusive execution category regardless of `execution_category`. + #[serde(default)] + pub terminal: bool, +} + +impl RuntimeToolDescriptor { + pub fn builder(name: impl Into) -> RuntimeToolDescriptorBuilder { + RuntimeToolDescriptorBuilder { + provider: ProviderToolSpec::builder(name), + capabilities: Vec::new(), + side_effect_level: ToolSideEffectLevel::None, + durability: ToolDurability::Ephemeral, + execution_category: ToolExecutionCategory::ExclusiveLocalMutation, + approval_category: ToolApprovalCategory::Default, + execution_timeout: None, + terminal: false, + } + } + + pub fn provider_spec(&self) -> &ProviderToolSpec { + &self.provider + } +} + +impl Deref for RuntimeToolDescriptor { + type Target = ProviderToolSpec; + + fn deref(&self) -> &Self::Target { + &self.provider + } +} + +#[derive(Debug, Clone)] +pub struct RuntimeToolDescriptorBuilder { + provider: ProviderToolSpecBuilder, + capabilities: Vec, + side_effect_level: ToolSideEffectLevel, + durability: ToolDurability, + execution_category: ToolExecutionCategory, + approval_category: ToolApprovalCategory, + execution_timeout: Option, + terminal: bool, +} + +impl RuntimeToolDescriptorBuilder { + pub fn description(mut self, description: impl Into) -> Self { + self.provider = self.provider.description(description); + self + } + + pub fn input_schema(mut self, input_schema: Value) -> Self { + self.provider = self.provider.input_schema(input_schema); + self + } + + pub fn output_schema(mut self, output_schema: Value) -> Self { + self.provider = self.provider.output_schema(output_schema); + self + } + + pub fn provider_kind(mut self, kind: ProviderToolKind) -> Self { + self.provider = self.provider.kind(kind); + self + } + + pub fn provider_options(mut self, options: Value) -> Self { + self.provider = self.provider.options(options); + self + } + + pub fn loading_policy(mut self, loading_policy: ToolLoadingPolicy) -> Self { + self.provider = self.provider.loading_policy(loading_policy); + self + } + + pub fn strict(mut self, strict: bool) -> Self { + self.provider = self.provider.strict(strict); + self + } + + pub fn non_strict(mut self) -> Self { + self.provider = self.provider.non_strict(); + self + } + + pub fn defer_loading(mut self, defer_loading: bool) -> Self { + self.provider = self.provider.defer_loading(defer_loading); + self + } + + pub fn capability(mut self, capability: ToolCapability) -> Self { + self.capabilities.push(capability); + self + } + + pub fn capabilities(mut self, capabilities: impl IntoIterator) -> Self { + self.capabilities = capabilities.into_iter().collect(); + self + } + + pub fn side_effect_level(mut self, side_effect_level: ToolSideEffectLevel) -> Self { + self.side_effect_level = side_effect_level; + self + } + + pub fn durability(mut self, durability: ToolDurability) -> Self { + self.durability = durability; + self + } + + pub fn execution_category(mut self, execution_category: ToolExecutionCategory) -> Self { + self.execution_category = execution_category; + self + } + + pub fn approval_category(mut self, approval_category: ToolApprovalCategory) -> Self { + self.approval_category = approval_category; + self + } + + pub fn execution_timeout(mut self, execution_timeout: Duration) -> Self { + self.execution_timeout = Some(execution_timeout); + self + } + + /// Marks this tool as a terminal action: a call is never scheduled in a + /// parallel batch regardless of `execution_category` — the scheduler + /// coerces it to an exclusive category, so it can never race parallel + /// retrieval calls scheduled in the same round. + pub fn terminal(mut self) -> Self { + self.terminal = true; + self + } + + pub fn build(self) -> RuntimeToolDescriptor { + RuntimeToolDescriptor { + provider: self.provider.build(), + capabilities: self.capabilities, + side_effect_level: self.side_effect_level, + durability: self.durability, + execution_category: self.execution_category, + approval_category: self.approval_category, + execution_timeout: self.execution_timeout, + terminal: self.terminal, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn descriptor_defaults_to_non_terminal() { + let descriptor = RuntimeToolDescriptor::builder("plain_tool").build(); + assert!(!descriptor.terminal); + } + + #[test] + fn terminal_builder_flag_marks_the_descriptor_terminal() { + let descriptor = RuntimeToolDescriptor::builder("finish_tool") + .terminal() + .build(); + assert!(descriptor.terminal); + } + + #[test] + fn terminal_flag_is_independent_of_declared_execution_category() { + let descriptor = RuntimeToolDescriptor::builder("finish_tool") + .execution_category(ToolExecutionCategory::ReadOnlyParallel) + .terminal() + .build(); + assert!(descriptor.terminal); + assert_eq!( + descriptor.execution_category, + ToolExecutionCategory::ReadOnlyParallel + ); + } +} diff --git a/vendor/mentra/src/tool/files.rs b/vendor/mentra/src/tool/files.rs new file mode 100644 index 0000000..021019d --- /dev/null +++ b/vendor/mentra/src/tool/files.rs @@ -0,0 +1,237 @@ +#[path = "files/execution.rs"] +mod execution; +#[path = "files/input.rs"] +mod input; +#[path = "files/preview.rs"] +mod preview; +#[path = "files/schema.rs"] +mod schema; +#[path = "files/workspace.rs"] +pub(crate) mod workspace; + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::tool::{ + ParallelToolContext, RuntimeToolDescriptor, ToolApprovalCategory, ToolAuthorizationPreview, + ToolCapability, ToolContext, ToolDefinition, ToolDurability, ToolExecutionCategory, + ToolExecutor, ToolResult, ToolSideEffectLevel, +}; + +use self::{ + execution::execute_files_tool, input::file_execution_category, + preview::build_files_authorization_preview, +}; + +pub struct FilesTool; + +impl ToolDefinition for FilesTool { + fn descriptor(&self) -> RuntimeToolDescriptor { + RuntimeToolDescriptor::builder("files") + .description("Read, search, list, create, update, move, and delete files within the workspace.") + .input_schema(json!({ + "type": "object", + "properties": { + "workingDirectory": { + "type": "string", + "description": "Optional directory used to resolve relative operation paths" + }, + "operations": { + "type": "array", + "description": "Ordered file operations to execute. Later reads can observe earlier staged writes.", + "items": { + "oneOf": [ + { + "type": "object", + "properties": { + "op": { "const": "read" }, + "path": { "type": "string" }, + "offset": { "type": "integer", "minimum": 1 }, + "limit": { "type": "integer", "minimum": 0 } + }, + "required": ["op", "path"] + }, + { + "type": "object", + "properties": { + "op": { "const": "list" }, + "path": { "type": "string" }, + "depth": { "type": "integer", "minimum": 0 }, + "limit": { "type": "integer", "minimum": 0 } + }, + "required": ["op", "path"] + }, + { + "type": "object", + "properties": { + "op": { "const": "search" }, + "path": { "type": "string" }, + "pattern": { "type": "string" }, + "limit": { "type": "integer", "minimum": 0 } + }, + "required": ["op", "path", "pattern"] + }, + { + "type": "object", + "properties": { + "op": { "const": "create" }, + "path": { "type": "string" }, + "content": { "type": "string" } + }, + "required": ["op", "path", "content"] + }, + { + "type": "object", + "properties": { + "op": { "const": "set" }, + "path": { "type": "string" }, + "content": { "type": "string" } + }, + "required": ["op", "path", "content"] + }, + { + "type": "object", + "properties": { + "op": { "const": "replace" }, + "path": { "type": "string" }, + "old": { "type": "string" }, + "new": { "type": "string" }, + "replaceAll": { "type": "boolean" }, + "expectedReplacements": { "type": "integer", "minimum": 0 } + }, + "required": ["op", "path", "old", "new"] + }, + { + "type": "object", + "properties": { + "op": { "const": "insert" }, + "path": { "type": "string" }, + "anchor": { "type": "string" }, + "position": { + "type": "string", + "enum": ["before", "after"] + }, + "content": { "type": "string" }, + "occurrence": { "type": "integer", "minimum": 1 } + }, + "required": ["op", "path", "anchor", "position", "content"] + }, + { + "type": "object", + "properties": { + "op": { "const": "move" }, + "from": { "type": "string" }, + "to": { "type": "string" } + }, + "required": ["op", "from", "to"] + }, + { + "type": "object", + "properties": { + "op": { "const": "delete" }, + "path": { "type": "string" } + }, + "required": ["op", "path"] + } + ] + } + } + }, + "required": ["operations"] + })) + .capabilities([ + ToolCapability::FilesystemRead, + ToolCapability::FilesystemWrite, + ]) + .side_effect_level(ToolSideEffectLevel::LocalState) + .durability(ToolDurability::Ephemeral) + .execution_category(ToolExecutionCategory::ExclusiveLocalMutation) + .approval_category(ToolApprovalCategory::Filesystem) + .build() + } +} + +#[async_trait] +impl ToolExecutor for FilesTool { + fn execution_category(&self, input: &Value) -> ToolExecutionCategory { + file_execution_category(input) + } + + fn authorization_preview( + &self, + ctx: &ParallelToolContext, + input: &Value, + ) -> Result { + build_files_authorization_preview(self.descriptor(), ctx, input) + } + + async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult { + execute_files_tool( + ctx.agent_id.clone(), + ctx.tool_call_id.clone(), + ctx.tool_name.clone(), + ctx.runtime.clone(), + ctx.resolve_working_directory(None) + .unwrap_or_else(|_| ctx.working_directory().to_path_buf()), + ctx.event_tx.clone(), + input, + ) + .await + } + + async fn execute_mut(&self, ctx: ToolContext<'_>, input: Value) -> ToolResult { + execute_files_tool( + ctx.agent_id.clone(), + ctx.tool_call_id.clone(), + ctx.tool_name.clone(), + ctx.runtime.clone(), + ctx.resolve_working_directory(None) + .unwrap_or_else(|_| ctx.working_directory().to_path_buf()), + ctx.event_tx.clone(), + input, + ) + .await + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn files_tool_metadata_marks_local_mutation() { + let spec = FilesTool.descriptor(); + assert!(!spec.capabilities.contains(&ToolCapability::ReadOnly)); + assert_eq!(spec.durability, ToolDurability::Ephemeral); + assert!(spec.capabilities.contains(&ToolCapability::FilesystemRead)); + assert!(spec.capabilities.contains(&ToolCapability::FilesystemWrite)); + assert_eq!( + spec.execution_category, + ToolExecutionCategory::ExclusiveLocalMutation + ); + } + + #[test] + fn read_only_operations_opt_into_parallel_execution() { + let category = FilesTool.execution_category(&json!({ + "operations": [ + { "op": "read", "path": "README.md" }, + { "op": "search", "path": ".", "pattern": "mentra" } + ] + })); + + assert_eq!(category, ToolExecutionCategory::ReadOnlyParallel); + } + + #[test] + fn mutating_operations_stay_exclusive() { + let category = FilesTool.execution_category(&json!({ + "operations": [ + { "op": "read", "path": "README.md" }, + { "op": "set", "path": "README.md", "content": "updated" } + ] + })); + + assert_eq!(category, ToolExecutionCategory::ExclusiveLocalMutation); + } +} diff --git a/vendor/mentra/src/tool/files/execution.rs b/vendor/mentra/src/tool/files/execution.rs new file mode 100644 index 0000000..54d3a6a --- /dev/null +++ b/vendor/mentra/src/tool/files/execution.rs @@ -0,0 +1,139 @@ +use serde_json::Value; + +use crate::{ + agent::{AgentEvent, AgentEventBus}, + runtime::RuntimeHandle, + tool::ToolResult, +}; + +use super::{ + input::{ensure_files_have_operations, parse_files_input}, + workspace::WorkspaceEditor, +}; + +pub(crate) async fn execute_files_tool( + agent_id: String, + tool_call_id: String, + tool_name: String, + runtime: RuntimeHandle, + default_working_directory: std::path::PathBuf, + event_tx: AgentEventBus, + input: Value, +) -> ToolResult { + let input = parse_files_input(&input)?; + ensure_files_have_operations(&input)?; + + let working_directory = match input.working_directory.as_deref() { + Some(directory) => runtime.resolve_working_directory(&agent_id, Some(directory))?, + None => runtime + .resolve_working_directory(&agent_id, None) + .unwrap_or(default_working_directory), + }; + let base_dir = runtime.agent_config(&agent_id)?.base_dir; + + tokio::task::spawn_blocking(move || { + let mut editor = WorkspaceEditor::new(agent_id, runtime, base_dir, working_directory); + let mut sections = Vec::with_capacity(input.operations.len()); + for operation in input.operations { + let section = editor.apply_operation(operation)?; + if let Some(progress) = file_op_progress(§ion) { + event_tx.send(AgentEvent::ToolExecutionProgress { + id: tool_call_id.clone(), + name: tool_name.clone(), + progress, + }); + } + sections.push(section); + } + editor.commit()?; + + Ok(sections.join("\n\n")) + }) + .await + .map_err(|error| format!("Files tool task failed: {error}"))? +} + +/// Derives a `file_op:` progress string from the operation summary returned by +/// `WorkspaceEditor::apply_operation`. Only mutating operations that produce a +/// recognisable prefix are surfaced; read-only operations return `None`. +fn file_op_progress(section: &str) -> Option { + // Mutating operation prefixes produced by WorkspaceEditor. + let mutating_prefixes = ["create ", "set ", "replace ", "insert ", "move ", "delete "]; + let first_line = section.lines().next().unwrap_or(section); + if mutating_prefixes + .iter() + .any(|prefix| first_line.starts_with(prefix)) + { + Some(format!("file_op: {first_line}")) + } else { + None + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn file_op_progress_create() { + let result = file_op_progress("create src/lib.rs"); + assert_eq!(result, Some("file_op: create src/lib.rs".to_string())); + } + + #[test] + fn file_op_progress_set() { + let result = file_op_progress("set src/main.rs"); + assert_eq!(result, Some("file_op: set src/main.rs".to_string())); + } + + #[test] + fn file_op_progress_replace() { + let result = file_op_progress("replace src/lib.rs (1 replacement)"); + assert_eq!( + result, + Some("file_op: replace src/lib.rs (1 replacement)".to_string()) + ); + } + + #[test] + fn file_op_progress_insert() { + let result = file_op_progress("insert src/lib.rs"); + assert_eq!(result, Some("file_op: insert src/lib.rs".to_string())); + } + + #[test] + fn file_op_progress_move() { + let result = file_op_progress("move old.rs -> new.rs"); + assert_eq!(result, Some("file_op: move old.rs -> new.rs".to_string())); + } + + #[test] + fn file_op_progress_delete() { + let result = file_op_progress("delete src/old.rs"); + assert_eq!(result, Some("file_op: delete src/old.rs".to_string())); + } + + #[test] + fn file_op_progress_read_returns_none() { + let result = file_op_progress("read src/lib.rs\nL1: fn main() {}"); + assert_eq!(result, None); + } + + #[test] + fn file_op_progress_list_returns_none() { + let result = file_op_progress("list src/\n[file] main.rs"); + assert_eq!(result, None); + } + + #[test] + fn file_op_progress_search_returns_none() { + let result = file_op_progress("search src/ /fn /\nsrc/lib.rs:1: fn main() {}"); + assert_eq!(result, None); + } + + #[test] + fn file_op_progress_uses_only_first_line() { + let result = file_op_progress("create foo.rs\nsome extra\ncontent"); + assert_eq!(result, Some("file_op: create foo.rs".to_string())); + } +} diff --git a/vendor/mentra/src/tool/files/input.rs b/vendor/mentra/src/tool/files/input.rs new file mode 100644 index 0000000..0cf7368 --- /dev/null +++ b/vendor/mentra/src/tool/files/input.rs @@ -0,0 +1,30 @@ +use serde_json::Value; + +use crate::tool::ToolExecutionCategory; + +use super::schema::{FileOperation, FilesInput}; + +pub(crate) fn parse_files_input(input: &Value) -> Result { + serde_json::from_value::(input.clone()) + .map_err(|error| format!("Invalid files input: {error}")) +} + +pub(crate) fn ensure_files_have_operations(input: &FilesInput) -> Result<(), String> { + if input.operations.is_empty() { + Err("At least one file operation is required".to_string()) + } else { + Ok(()) + } +} + +pub(crate) fn file_execution_category(input: &Value) -> ToolExecutionCategory { + let Ok(input) = parse_files_input(input) else { + return ToolExecutionCategory::ExclusiveLocalMutation; + }; + + if input.operations.iter().all(FileOperation::is_read_only) { + ToolExecutionCategory::ReadOnlyParallel + } else { + ToolExecutionCategory::ExclusiveLocalMutation + } +} diff --git a/vendor/mentra/src/tool/files/preview.rs b/vendor/mentra/src/tool/files/preview.rs new file mode 100644 index 0000000..b8ceda9 --- /dev/null +++ b/vendor/mentra/src/tool/files/preview.rs @@ -0,0 +1,171 @@ +use serde_json::{Value, json}; +use std::path::{Component, Path, PathBuf}; + +use crate::tool::{ParallelToolContext, RuntimeToolDescriptor, ToolAuthorizationPreview}; + +use super::{ + input::{ensure_files_have_operations, parse_files_input}, + schema::{self, FileOperation}, +}; + +pub(crate) fn build_files_authorization_preview( + descriptor: RuntimeToolDescriptor, + ctx: &ParallelToolContext, + input: &Value, +) -> Result { + let raw_input = input.clone(); + let input = parse_files_input(input)?; + ensure_files_have_operations(&input)?; + + let working_directory = match input.working_directory.as_deref() { + Some(directory) => ctx.resolve_working_directory(Some(directory))?, + None => ctx + .runtime + .resolve_working_directory(&ctx.agent_id, None) + .unwrap_or_else(|_| ctx.working_directory().to_path_buf()), + }; + + let operations = input + .operations + .into_iter() + .map(|operation| preview_file_operation(&working_directory, operation)) + .collect::, _>>()?; + + Ok(ToolAuthorizationPreview { + working_directory: working_directory.clone(), + capabilities: descriptor.capabilities, + side_effect_level: descriptor.side_effect_level, + durability: descriptor.durability, + execution_category: descriptor.execution_category, + approval_category: descriptor.approval_category, + raw_input, + structured_input: json!({ + "kind": "files", + "working_directory": working_directory, + "operations": operations, + }), + }) +} + +fn preview_file_operation( + working_directory: &Path, + operation: FileOperation, +) -> Result { + match operation { + FileOperation::Read { + path, + offset, + limit, + } => Ok(json!({ + "op": "read", + "resolved_path": resolve_preview_path(working_directory, &path)?, + "offset": offset, + "limit": limit, + })), + FileOperation::List { path, depth, limit } => Ok(json!({ + "op": "list", + "resolved_path": resolve_preview_path(working_directory, &path)?, + "depth": depth, + "limit": limit, + })), + FileOperation::Search { + path, + pattern, + limit, + } => Ok(json!({ + "op": "search", + "resolved_path": resolve_preview_path(working_directory, &path)?, + "pattern": pattern, + "limit": limit, + })), + FileOperation::Create { path, .. } => Ok(json!({ + "op": "create", + "resolved_path": resolve_preview_path(working_directory, &path)?, + })), + FileOperation::Set { path, .. } => Ok(json!({ + "op": "set", + "resolved_path": resolve_preview_path(working_directory, &path)?, + })), + FileOperation::Replace { + path, + replace_all, + expected_replacements, + .. + } => Ok(json!({ + "op": "replace", + "resolved_path": resolve_preview_path(working_directory, &path)?, + "replace_all": replace_all, + "expected_replacements": expected_replacements, + })), + FileOperation::Insert { + path, + position, + occurrence, + .. + } => Ok(json!({ + "op": "insert", + "resolved_path": resolve_preview_path(working_directory, &path)?, + "position": match position { + schema::InsertPosition::Before => "before", + schema::InsertPosition::After => "after", + }, + "occurrence": occurrence, + })), + FileOperation::Move { from, to } => Ok(json!({ + "op": "move", + "from_resolved_path": resolve_preview_path(working_directory, &from)?, + "to_resolved_path": resolve_preview_path(working_directory, &to)?, + })), + FileOperation::Delete { path } => Ok(json!({ + "op": "delete", + "resolved_path": resolve_preview_path(working_directory, &path)?, + })), + } +} + +fn resolve_preview_path(working_directory: &Path, raw: &str) -> Result { + let candidate = PathBuf::from(raw); + let path = if candidate.is_absolute() { + candidate + } else { + working_directory.join(candidate) + }; + normalize_preview_path(path) +} + +fn normalize_preview_path(path: PathBuf) -> Result { + let mut normalized = if path.is_absolute() { + PathBuf::new() + } else { + return Err(format!( + "Path '{}' must resolve to an absolute path", + path.display() + )); + }; + + for component in path.components() { + match component { + Component::Prefix(prefix) => normalized.push(prefix.as_os_str()), + Component::RootDir => normalized.push(component.as_os_str()), + Component::CurDir => {} + Component::ParentDir => { + if !normalized.pop() || !normalized.is_absolute() { + return Err(format!( + "Path '{}' escapes the filesystem root", + path.display() + )); + } + } + Component::Normal(segment) => normalized.push(segment), + } + } + + if !normalized.is_absolute() { + return Err(format!( + "Path '{}' must resolve to an absolute path", + path.display() + )); + } + + Ok(normalized) +} diff --git a/vendor/mentra/src/tool/files/schema.rs b/vendor/mentra/src/tool/files/schema.rs new file mode 100644 index 0000000..b0cf8ae --- /dev/null +++ b/vendor/mentra/src/tool/files/schema.rs @@ -0,0 +1,75 @@ +use serde::Deserialize; + +#[derive(Debug, Deserialize)] +pub(crate) struct FilesInput { + #[serde(rename = "workingDirectory")] + pub(crate) working_directory: Option, + pub(crate) operations: Vec, +} + +#[derive(Debug, Deserialize)] +#[serde(tag = "op", rename_all = "snake_case")] +pub(crate) enum FileOperation { + Read { + path: String, + offset: Option, + limit: Option, + }, + List { + path: String, + depth: Option, + limit: Option, + }, + Search { + path: String, + pattern: String, + limit: Option, + }, + Create { + path: String, + content: String, + }, + Set { + path: String, + content: String, + }, + Replace { + path: String, + old: String, + new: String, + #[serde(rename = "replaceAll")] + replace_all: Option, + #[serde(rename = "expectedReplacements")] + expected_replacements: Option, + }, + Insert { + path: String, + anchor: String, + position: InsertPosition, + content: String, + occurrence: Option, + }, + Move { + from: String, + to: String, + }, + Delete { + path: String, + }, +} + +impl FileOperation { + pub(crate) fn is_read_only(&self) -> bool { + matches!( + self, + Self::Read { .. } | Self::List { .. } | Self::Search { .. } + ) + } +} + +#[derive(Debug, Clone, Copy, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum InsertPosition { + Before, + After, +} diff --git a/vendor/mentra/src/tool/files/workspace.rs b/vendor/mentra/src/tool/files/workspace.rs new file mode 100644 index 0000000..b29c01c --- /dev/null +++ b/vendor/mentra/src/tool/files/workspace.rs @@ -0,0 +1,1006 @@ +#[path = "workspace/edit.rs"] +mod edit; +#[path = "workspace/operations.rs"] +mod operations; + +pub(crate) use edit::{EditOutcome, TextEdit}; + +use std::{ + collections::{BTreeMap, BTreeSet}, + fs, + path::{Component, Path, PathBuf}, + time::{SystemTime, UNIX_EPOCH}, +}; + +use regex::Regex; + +use crate::runtime::{RuntimeHandle, RuntimeHookEvent}; + +#[derive(Debug, Clone)] +enum OverlayEntry { + File(Vec), + Deleted, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum EntryKind { + File, + Dir, + Missing, +} + +#[derive(Debug, Clone)] +enum OriginalState { + Missing, + File(Vec), +} + +#[derive(Debug, Clone, Default)] +pub(crate) struct SearchOptions { + pub(crate) file_glob: Option, + pub(crate) ignore_case: bool, + pub(crate) literal: bool, + pub(crate) context: usize, + pub(crate) multiline: bool, + pub(crate) max_line_chars: Option, +} + +pub(crate) struct WorkspaceEditor { + agent_id: String, + runtime: RuntimeHandle, + base_dir: PathBuf, + working_directory: PathBuf, + overlay: BTreeMap, +} + +impl WorkspaceEditor { + pub(crate) fn new( + agent_id: String, + runtime: RuntimeHandle, + base_dir: PathBuf, + working_directory: PathBuf, + ) -> Self { + let base_dir = canonicalize_existing_path(base_dir); + let working_directory = canonicalize_existing_path(working_directory); + Self { + agent_id, + runtime, + base_dir, + working_directory, + overlay: BTreeMap::new(), + } + } + + pub(crate) fn commit(&self) -> Result<(), String> { + if self.overlay.is_empty() { + return Ok(()); + } + + let originals = self + .overlay + .keys() + .map(|path| Ok((path.clone(), self.capture_original_state(path)?))) + .collect::, String>>()?; + + let file_writes = self + .overlay + .iter() + .filter_map(|(path, entry)| match entry { + OverlayEntry::File(bytes) => Some((path, bytes.as_slice())), + OverlayEntry::Deleted => None, + }) + .collect::>(); + let deletes = self + .overlay + .iter() + .filter_map(|(path, entry)| match entry { + OverlayEntry::Deleted => Some(path), + OverlayEntry::File(_) => None, + }) + .collect::>(); + + let result = (|| -> Result<(), String> { + for (path, bytes) in file_writes { + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).map_err(|error| { + format!("Failed to create directory '{}': {error}", parent.display()) + })?; + } + let temp_path = temporary_path(path); + fs::write(&temp_path, bytes).map_err(|error| { + format!("Failed to write '{}': {error}", temp_path.display()) + })?; + replace_file(&temp_path, path)?; + } + + for path in deletes { + if path.exists() { + fs::remove_file(path).map_err(|error| { + format!("Failed to delete '{}': {error}", path.display()) + })?; + } + } + + Ok(()) + })(); + + if let Err(error) = result { + let _ = self.rollback(&originals); + return Err(error); + } + + Ok(()) + } + + fn rollback(&self, originals: &BTreeMap) -> Result<(), String> { + for (path, original) in originals { + match original { + OriginalState::Missing => { + if path.exists() { + fs::remove_file(path).map_err(|error| { + format!("Failed to roll back '{}': {error}", path.display()) + })?; + } + } + OriginalState::File(bytes) => { + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).map_err(|error| { + format!( + "Failed to recreate directory '{}' during rollback: {error}", + parent.display() + ) + })?; + } + let temp_path = temporary_path(path); + fs::write(&temp_path, bytes).map_err(|error| { + format!( + "Failed to write rollback temp file '{}': {error}", + temp_path.display() + ) + })?; + replace_file(&temp_path, path)?; + } + } + } + Ok(()) + } + + fn resolve_path(&self, raw: &str) -> Result { + let candidate = PathBuf::from(raw); + let path = if candidate.is_absolute() { + candidate + } else { + self.working_directory.join(candidate) + }; + normalize_path(path) + } + + fn authorize_read(&self, path: &Path, action: &str) -> Result { + self.runtime + .execution + .policy + .authorize_file_read(&self.base_dir, path) + .inspect_err(|detail: &String| { + let _ = self + .runtime + .emit_hook(RuntimeHookEvent::AuthorizationDenied { + agent_id: self.agent_id.clone(), + action: action.to_string(), + detail: detail.clone(), + }); + }) + } + + fn authorize_write(&self, path: &Path, action: &str) -> Result { + self.runtime + .execution + .policy + .authorize_file_write(&self.base_dir, path) + .inspect_err(|detail: &String| { + let _ = self + .runtime + .emit_hook(RuntimeHookEvent::AuthorizationDenied { + agent_id: self.agent_id.clone(), + action: action.to_string(), + detail: detail.clone(), + }); + }) + } + + fn entry_kind(&self, path: &Path) -> Result { + if let Some(entry) = self.overlay.get(path) { + return Ok(match entry { + OverlayEntry::File(_) => EntryKind::File, + OverlayEntry::Deleted => EntryKind::Missing, + }); + } + + if self.has_live_descendant(path) { + return Ok(EntryKind::Dir); + } + + match fs::metadata(path) { + Ok(metadata) if metadata.is_dir() => Ok(EntryKind::Dir), + Ok(metadata) if metadata.is_file() => Ok(EntryKind::File), + Ok(_) => Err(format!( + "Path '{}' is not a regular file or directory", + self.display_path(path) + )), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(EntryKind::Missing), + Err(error) => Err(format!( + "Failed to inspect '{}': {error}", + self.display_path(path) + )), + } + } + + fn has_live_descendant(&self, path: &Path) -> bool { + self.overlay.iter().any(|(candidate, entry)| match entry { + OverlayEntry::File(_) => candidate.starts_with(path) && candidate != path, + OverlayEntry::Deleted => false, + }) + } + + fn child_names(&self, dir: &Path) -> Result, String> { + let mut names = BTreeSet::new(); + + match fs::read_dir(dir) { + Ok(entries) => { + for entry in entries { + let entry = entry.map_err(|error| { + format!( + "Failed to read directory '{}': {error}", + self.display_path(dir) + ) + })?; + names.insert(entry.file_name().to_string_lossy().into_owned()); + } + } + Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} + Err(error) => { + return Err(format!( + "Failed to read directory '{}': {error}", + self.display_path(dir) + )); + } + } + + for path in self.overlay.keys() { + if let Ok(relative) = path.strip_prefix(dir) + && let Some(Component::Normal(name)) = relative.components().next() + { + names.insert(name.to_string_lossy().into_owned()); + } + } + + Ok(names.into_iter().collect()) + } + + fn walk_entries( + &self, + dir: &Path, + current_depth: usize, + max_depth: usize, + action: &str, + visited: &mut BTreeSet, + visit: &mut F, + ) -> Result + where + F: FnMut(&Path, EntryKind) -> Result, + { + if current_depth > max_depth { + return Ok(true); + } + if !self.mark_directory_visited(dir, visited)? { + return Ok(true); + } + + for child_name in self.child_names(dir)? { + let child = dir.join(&child_name); + // The traversal root was authorized before recursion starts, but a + // descendant may be a symlink whose target leaves every allowed + // read root. Reauthorize each child immediately before inspecting + // or following it so recursion cannot cross that boundary. Doing + // this lazily also preserves the existing limit short-circuit. + self.authorize_read(&child, action)?; + let kind = self.entry_kind(&child)?; + if kind == EntryKind::Missing { + continue; + } + if !visit(&child, kind)? { + return Ok(false); + } + if kind == EntryKind::Dir + && !self.walk_entries( + &child, + current_depth + 1, + max_depth, + action, + visited, + visit, + )? + { + return Ok(false); + } + } + + Ok(true) + } + + fn collect_search_matches( + &self, + root: &Path, + regex: &Regex, + options: &SearchOptions, + limit: usize, + matches: &mut Vec, + ) -> Result<(), String> { + let mut visited = BTreeSet::new(); + self.walk_entries( + root, + 1, + usize::MAX, + "files_search", + &mut visited, + &mut |child, kind| { + if kind == EntryKind::File + && self.path_matches_glob(root, child, options.file_glob.as_deref()) + { + self.search_file(child, regex, options, limit, matches)?; + } + Ok(matches.len() < limit) + }, + )?; + Ok(()) + } + + fn collect_glob_matches( + &self, + root: &Path, + pattern: &str, + limit: usize, + matches: &mut Vec, + ) -> Result<(), String> { + let mut visited = BTreeSet::new(); + self.walk_entries( + root, + 1, + usize::MAX, + "files_glob", + &mut visited, + &mut |child, kind| { + if kind == EntryKind::File && self.path_matches_glob(root, child, Some(pattern)) { + matches.push(self.display_relative_to(root, child)); + } + Ok(matches.len() < limit) + }, + )?; + Ok(()) + } + + fn mark_directory_visited( + &self, + dir: &Path, + visited: &mut BTreeSet, + ) -> Result { + let key = match fs::canonicalize(dir) { + Ok(path) => path, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => dir.to_path_buf(), + Err(error) => { + return Err(format!( + "Failed to resolve directory '{}': {error}", + self.display_path(dir) + )); + } + }; + Ok(visited.insert(key)) + } + + fn search_file( + &self, + path: &Path, + regex: &Regex, + options: &SearchOptions, + limit: usize, + matches: &mut Vec, + ) -> Result<(), String> { + let content = self.load_text_file(path)?; + let lines = content.lines().collect::>(); + if lines.is_empty() { + return Ok(()); + } + let mut matched_lines = BTreeSet::new(); + + if options.multiline { + let line_starts = line_start_offsets(&content); + let last_line = lines.len() - 1; + for found in regex.find_iter(&content) { + let start_line = line_index_at(&line_starts, found.start()).min(last_line); + let end_position = found.end().saturating_sub(1).max(found.start()); + let end_line = line_index_at(&line_starts, end_position).min(last_line); + matched_lines.extend(start_line..=end_line); + } + } else { + matched_lines.extend( + lines + .iter() + .enumerate() + .filter_map(|(index, line)| regex.is_match(line).then_some(index)), + ); + } + + let mut rendered_lines = BTreeSet::new(); + for index in &matched_lines { + let start = index.saturating_sub(options.context); + let end = index + .saturating_add(options.context) + .saturating_add(1) + .min(lines.len()); + rendered_lines.extend(start..end); + } + + for index in rendered_lines { + if matches.len() >= limit { + break; + } + let separator = if matched_lines.contains(&index) { + ':' + } else { + '-' + }; + matches.push(format!( + "{}:{}{} {}", + self.display_path(path), + index + 1, + separator, + options.max_line_chars.map_or_else( + || lines[index].to_string(), + |max_chars| cap_search_line(lines[index], max_chars), + ) + )); + } + Ok(()) + } + + /// Whether a workspace file matches a caller's glob. + /// + /// A path glob, deliberately: the pattern describes a filesystem path, so + /// `/` is a separator that `*` must not cross and `**/*.rs` has to mean + /// any depth. This is the opposite of a permission rule pattern, which + /// describes JSON and uses [`crate::session::permission`]'s separatorless + /// matcher instead. + fn path_matches_glob(&self, root: &Path, path: &Path, pattern: Option<&str>) -> bool { + let Some(pattern) = pattern else { + return true; + }; + let relative = self.display_relative_to(root, path); + glob_match::glob_match(pattern, &relative) + || (!pattern.contains('/') + && path + .file_name() + .is_some_and(|name| glob_match::glob_match(pattern, &name.to_string_lossy()))) + } + + fn load_text_file(&self, path: &Path) -> Result { + let bytes = self.load_file_bytes(path)?; + String::from_utf8(bytes) + .map_err(|_| format!("Path '{}' is not valid UTF-8 text", self.display_path(path))) + } + + fn load_file_bytes(&self, path: &Path) -> Result, String> { + if let Some(entry) = self.overlay.get(path) { + return match entry { + OverlayEntry::File(bytes) => Ok(bytes.clone()), + OverlayEntry::Deleted => { + Err(format!("Path '{}' does not exist", self.display_path(path))) + } + }; + } + + match fs::read(path) { + Ok(bytes) => Ok(bytes), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + Err(format!("Path '{}' does not exist", self.display_path(path))) + } + Err(error) => Err(format!( + "Failed to read '{}': {error}", + self.display_path(path) + )), + } + } + + fn capture_original_state(&self, path: &Path) -> Result { + match fs::read(path) { + Ok(bytes) => Ok(OriginalState::File(bytes)), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + Ok(OriginalState::Missing) + } + Err(error) => Err(format!( + "Failed to snapshot '{}': {error}", + self.display_path(path) + )), + } + } + + fn display_path(&self, path: &Path) -> String { + if let Ok(relative) = path.strip_prefix(&self.working_directory) { + let rendered = relative.display().to_string(); + if rendered.is_empty() { + ".".to_string() + } else { + normalize_display_path(rendered) + } + } else { + normalize_display_path(path.display().to_string()) + } + } + + fn display_relative_to(&self, root: &Path, path: &Path) -> String { + if path == root { + path.file_name() + .map(|name| name.to_string_lossy().into_owned()) + .unwrap_or_else(|| ".".to_string()) + } else { + path.strip_prefix(root) + .map(|relative| normalize_display_path(relative.display().to_string())) + .unwrap_or_else(|_| self.display_path(path)) + } + } +} + +fn line_start_offsets(content: &str) -> Vec { + let mut offsets = vec![0]; + offsets.extend( + content + .bytes() + .enumerate() + .filter_map(|(index, byte)| (byte == b'\n').then_some(index + 1)), + ); + offsets +} + +fn line_index_at(line_starts: &[usize], byte_offset: usize) -> usize { + line_starts + .partition_point(|start| *start <= byte_offset) + .saturating_sub(1) +} + +fn cap_search_line(line: &str, max_chars: usize) -> String { + if max_chars == 0 { + return String::new(); + } + let mut chars = line.chars(); + let retained = chars.by_ref().take(max_chars).collect::(); + if chars.next().is_none() { + retained + } else { + let mut capped = retained.chars().take(max_chars - 1).collect::(); + capped.push('…'); + capped + } +} + +fn normalize_display_path(path: String) -> String { + path.replace('\\', "/") +} + +fn normalize_path(path: PathBuf) -> Result { + let mut normalized = if path.is_absolute() { + PathBuf::new() + } else { + return Err(format!( + "Path '{}' must resolve to an absolute path", + path.display() + )); + }; + + for component in path.components() { + match component { + Component::Prefix(prefix) => normalized.push(prefix.as_os_str()), + Component::RootDir => normalized.push(component.as_os_str()), + Component::CurDir => {} + Component::ParentDir => { + if !normalized.pop() || !normalized.is_absolute() { + return Err(format!( + "Path '{}' escapes the filesystem root", + path.display() + )); + } + } + Component::Normal(segment) => normalized.push(segment), + } + } + + if !normalized.is_absolute() { + return Err(format!( + "Path '{}' must resolve to an absolute path", + path.display() + )); + } + + Ok(normalized) +} + +fn temporary_path(path: &Path) -> PathBuf { + let unique = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or(0); + let file_name = path + .file_name() + .map(|name| name.to_string_lossy().into_owned()) + .unwrap_or_else(|| "file".to_string()); + path.with_file_name(format!(".{file_name}.mentra-tmp-{unique}")) +} + +fn canonicalize_existing_path(path: PathBuf) -> PathBuf { + fs::canonicalize(&path).unwrap_or(path) +} + +fn replace_file(temp_path: &Path, path: &Path) -> Result<(), String> { + #[cfg(windows)] + if path.exists() { + fs::remove_file(path) + .map_err(|error| format!("Failed to replace existing '{}': {error}", path.display()))?; + } + + fs::rename(temp_path, path).map_err(|error| { + format!( + "Failed to rename '{}' into '{}': {error}", + temp_path.display(), + path.display() + ) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn test_editor(label: &str) -> (PathBuf, WorkspaceEditor) { + let unique = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("duration") + .as_nanos(); + let root = std::env::temp_dir().join(format!("mentra-workspace-{label}-{unique}")); + fs::create_dir_all(&root).expect("create test workspace"); + let editor = WorkspaceEditor::new( + "agent".to_string(), + RuntimeHandle::new(false), + root.clone(), + root.clone(), + ); + (root, editor) + } + + #[test] + fn normalize_path_rejects_parent_past_root() { + let mut path = std::env::temp_dir(); + for _ in 0..10 { + path.push(".."); + } + path.push("escape"); + let error = normalize_path(path).expect_err("path should be rejected"); + assert!(error.contains("escapes the filesystem root")); + } + + #[test] + fn glob_walks_nested_files_with_workspace_relative_patterns() { + let (root, editor) = test_editor("glob"); + fs::create_dir_all(root.join("src/nested")).expect("create nested directory"); + fs::write(root.join("src/lib.rs"), "pub fn lib() {}\n").expect("write lib"); + fs::write(root.join("src/nested/mod.rs"), "pub mod nested;\n").expect("write module"); + fs::write(root.join("src/note.txt"), "not rust\n").expect("write note"); + + let output = editor.glob(".".to_string(), "**/*.rs", 20).expect("glob"); + + assert!(output.contains("src/lib.rs")); + assert!(output.contains("src/nested/mod.rs")); + assert!(!output.contains("note.txt")); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[cfg(unix)] + #[test] + fn recursive_limits_short_circuit_before_unvisited_descendants() { + use std::os::unix::fs::symlink; + + let (root, editor) = test_editor("walk-limit"); + let outside = root.with_file_name(format!( + "{}-outside", + root.file_name().expect("root name").to_string_lossy() + )); + fs::create_dir_all(&outside).expect("create outside directory"); + fs::write(root.join("a.txt"), "match\n").expect("write first file"); + symlink(&outside, root.join("z_escape")).expect("create escape symlink"); + + let list = editor.list(".".to_string(), 1, 1).expect("limited list"); + assert!(list.contains("[file] a.txt")); + let search = editor + .search(".".to_string(), "match", 1) + .expect("limited search"); + assert!(search.contains("a.txt:1: match")); + + fs::remove_dir_all(root).expect("remove test workspace"); + fs::remove_dir_all(outside).expect("remove outside directory"); + } + + #[test] + fn grep_supports_multiline_regex_and_context() { + let (root, editor) = test_editor("multiline-grep"); + fs::write( + root.join("multi.txt"), + "before\nBEGIN\nmiddle\nEND\nafter\n", + ) + .expect("write multiline file"); + + let output = editor + .grep( + "multi.txt".to_string(), + "BEGIN.*END", + SearchOptions { + multiline: true, + context: 1, + ..Default::default() + }, + 20, + ) + .expect("grep"); + + assert!(output.contains("multi.txt:1- before")); + assert!(output.contains("multi.txt:2: BEGIN")); + assert!(output.contains("multi.txt:3: middle")); + assert!(output.contains("multi.txt:4: END")); + assert!(output.contains("multi.txt:5- after")); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn grep_caps_each_physical_line_at_500_unicode_characters() { + let (root, editor) = test_editor("grep-line-cap"); + fs::write(root.join("long.txt"), format!("{}\n", "界".repeat(600))) + .expect("write long line"); + + let output = editor + .grep( + "long.txt".to_string(), + "界", + SearchOptions { + max_line_chars: Some(500), + ..Default::default() + }, + 20, + ) + .expect("grep"); + let rendered = output + .lines() + .find_map(|line| line.strip_prefix("long.txt:1: ")) + .expect("rendered match"); + + assert_eq!(rendered.chars().count(), 500); + assert!(rendered.ends_with('…')); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn legacy_batched_search_keeps_uncapped_matching_lines() { + let (root, editor) = test_editor("legacy-search-line"); + fs::write(root.join("long.txt"), format!("{}\n", "界".repeat(600))) + .expect("write long line"); + + let output = editor + .search("long.txt".to_string(), "界", 20) + .expect("search"); + let rendered = output + .lines() + .find_map(|line| line.strip_prefix("long.txt:1: ")) + .expect("rendered match"); + + assert_eq!(rendered.chars().count(), 600); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn multiline_grep_bounds_zero_width_match_at_trailing_newline() { + let (root, editor) = test_editor("grep-zero-width"); + fs::write(root.join("line.txt"), "line\n").expect("write line"); + + let output = editor + .grep( + "line.txt".to_string(), + "$", + SearchOptions { + multiline: true, + ..Default::default() + }, + 20, + ) + .expect("grep"); + + assert!(output.contains("line.txt:1: line")); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn grep_combines_literal_case_insensitive_and_file_glob_options() { + let (root, editor) = test_editor("grep-options"); + fs::write(root.join("code.rs"), "Needle.[x]\n").expect("write Rust file"); + fs::write(root.join("note.txt"), "Needle.[x]\n").expect("write text file"); + + let output = editor + .grep( + ".".to_string(), + "needle.[X]", + SearchOptions { + file_glob: Some("*.rs".to_string()), + ignore_case: true, + literal: true, + ..Default::default() + }, + 20, + ) + .expect("grep"); + + assert!(output.contains("code.rs:1: Needle.[x]")); + assert!(!output.contains("note.txt")); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn edit_restores_bom_and_crlf_and_reports_diff_metadata() { + let (root, mut editor) = test_editor("edit-crlf"); + fs::write( + root.join("note.txt"), + b"\xEF\xBB\xBFfirst\r\nsecond\r\nthird\r\n", + ) + .expect("write CRLF document"); + + let outcome = editor + .edit( + "note.txt".to_string(), + vec![TextEdit { + old_string: "second\r\n".to_string(), + new_string: "changed\r\n".to_string(), + }], + false, + ) + .expect("edit"); + editor.commit().expect("commit edit"); + + assert_eq!(outcome.replacement_count, 1); + assert_eq!(outcome.first_changed_line, 2); + assert!(outcome.diff.contains("-second")); + assert!(outcome.diff.contains("+changed")); + assert!(outcome.patch.contains("--- note.txt")); + assert_eq!( + fs::read(root.join("note.txt")).expect("read edited document"), + b"\xEF\xBB\xBFfirst\r\nchanged\r\nthird\r\n" + ); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn fuzzy_edit_normalizes_unicode_but_preserves_unchanged_original_lines() { + let (root, mut editor) = test_editor("edit-fuzzy"); + fs::write( + root.join("note.txt"), + "let label = “A—value”; \nlet count = 1;\n", + ) + .expect("write fuzzy document"); + + editor + .edit( + "note.txt".to_string(), + vec![TextEdit { + old_string: "let label = \"A-value\";\nlet count = 1;".to_string(), + new_string: "let label = \"A-value\";\nlet count = 2;".to_string(), + }], + false, + ) + .expect("fuzzy edit"); + editor.commit().expect("commit edit"); + + assert_eq!( + fs::read_to_string(root.join("note.txt")).expect("read edited document"), + "let label = “A—value”; \nlet count = 2;\n" + ); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn multi_edit_matches_against_original_content_and_rejects_overlap() { + let (root, mut editor) = test_editor("edit-original"); + fs::write(root.join("note.txt"), "alpha beta\n").expect("write document"); + + editor + .edit( + "note.txt".to_string(), + vec![ + TextEdit { + old_string: "alpha".to_string(), + new_string: "beta".to_string(), + }, + TextEdit { + old_string: "beta".to_string(), + new_string: "gamma".to_string(), + }, + ], + false, + ) + .expect("non-overlapping edit"); + editor.commit().expect("commit edit"); + assert_eq!( + fs::read_to_string(root.join("note.txt")).expect("read edited document"), + "beta gamma\n" + ); + + fs::write(root.join("overlap.txt"), "abcdef\n").expect("write overlap document"); + let error = editor + .edit( + "overlap.txt".to_string(), + vec![ + TextEdit { + old_string: "abcd".to_string(), + new_string: "one".to_string(), + }, + TextEdit { + old_string: "cdef".to_string(), + new_string: "two".to_string(), + }, + ], + false, + ) + .expect_err("overlap must fail"); + assert!(error.contains("overlap")); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn edit_rejects_ambiguous_and_no_op_replacements() { + let (root, mut editor) = test_editor("edit-guards"); + fs::write(root.join("note.txt"), "same same\n").expect("write document"); + + let ambiguous = editor + .edit( + "note.txt".to_string(), + vec![TextEdit { + old_string: "same".to_string(), + new_string: "changed".to_string(), + }], + false, + ) + .expect_err("ambiguous edit must fail"); + assert!(ambiguous.contains("not unique")); + + let no_op = editor + .edit( + "note.txt".to_string(), + vec![TextEdit { + old_string: "same".to_string(), + new_string: "same".to_string(), + }], + true, + ) + .expect_err("no-op edit must fail"); + assert!(no_op.contains("no-op")); + fs::remove_dir_all(root).expect("remove test workspace"); + } + + #[test] + fn legacy_batched_replace_keeps_its_first_match_contract() { + let (root, mut editor) = test_editor("legacy-replace"); + fs::write(root.join("note.txt"), "same same\n").expect("write document"); + + let summary = editor + .replace("note.txt".to_string(), "same", "changed", false, 2) + .expect("legacy replace"); + editor.commit().expect("commit replace"); + + assert_eq!(summary, "replace note.txt (2 replacements)"); + assert_eq!( + fs::read_to_string(root.join("note.txt")).expect("read document"), + "changed same\n" + ); + fs::remove_dir_all(root).expect("remove test workspace"); + } +} diff --git a/vendor/mentra/src/tool/files/workspace/edit.rs b/vendor/mentra/src/tool/files/workspace/edit.rs new file mode 100644 index 0000000..e58b24e --- /dev/null +++ b/vendor/mentra/src/tool/files/workspace/edit.rs @@ -0,0 +1,386 @@ +use similar::{ChangeTag, TextDiff}; +use unicode_normalization::UnicodeNormalization; + +use super::{EntryKind, OverlayEntry, WorkspaceEditor}; + +const UTF8_BOM: &[u8] = b"\xEF\xBB\xBF"; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct TextEdit { + pub(crate) old_string: String, + pub(crate) new_string: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct EditOutcome { + pub(crate) display_path: String, + pub(crate) replacement_count: usize, + pub(crate) diff: String, + pub(crate) patch: String, + pub(crate) first_changed_line: usize, +} + +#[derive(Debug, Clone)] +struct PlannedReplacement { + start: usize, + end: usize, + replacement: String, + edit_index: usize, +} + +#[derive(Debug, Clone, Copy)] +struct LineSpan<'a> { + start: usize, + content_end: usize, + full_end: usize, + text: &'a str, +} + +#[derive(Debug, Clone, Copy)] +enum LineEnding { + Lf, + CrLf, +} + +impl WorkspaceEditor { + pub(crate) fn edit( + &mut self, + path: String, + edits: Vec, + replace_all: bool, + ) -> Result { + if edits.is_empty() { + return Err("At least one edit is required".to_string()); + } + + let path = self.resolve_path(&path)?; + let path = self.authorize_write(&path, "files_write")?; + if self.entry_kind(&path)? != EntryKind::File { + return Err(format!( + "Path '{}' does not exist as a file", + self.display_path(&path) + )); + } + + let bytes = self.load_file_bytes(&path)?; + let has_bom = bytes.starts_with(UTF8_BOM); + let text_bytes = if has_bom { + &bytes[UTF8_BOM.len()..] + } else { + &bytes + }; + let source = std::str::from_utf8(text_bytes).map_err(|_| { + format!( + "Path '{}' is not valid UTF-8 text", + self.display_path(&path) + ) + })?; + let line_ending = if source.contains("\r\n") { + LineEnding::CrLf + } else { + LineEnding::Lf + }; + let original = source.replace("\r\n", "\n"); + + let edits = edits + .into_iter() + .map(|edit| TextEdit { + old_string: normalize_edit_text(&edit.old_string), + new_string: normalize_edit_text(&edit.new_string), + }) + .collect::>(); + let mut replacements = Vec::new(); + for (edit_index, edit) in edits.iter().enumerate() { + validate_edit(edit, edit_index)?; + let mut ranges = exact_ranges(&original, &edit.old_string); + let fuzzy = ranges.is_empty(); + if fuzzy { + ranges = fuzzy_ranges(&original, &edit.old_string); + } + + if ranges.is_empty() { + return Err(format!( + "Edit {} old_string was not found in '{}'", + edit_index + 1, + self.display_path(&path) + )); + } + if !replace_all && ranges.len() != 1 { + return Err(format!( + "Edit {} old_string is not unique in '{}' ({} matches); set replace_all to replace every match", + edit_index + 1, + self.display_path(&path), + ranges.len() + )); + } + + let selected = if replace_all { + ranges.as_slice() + } else { + &ranges[..1] + }; + for &(start, end) in selected { + let replacement = if fuzzy { + overlay_fuzzy_replacement( + &original[start..end], + &edit.old_string, + &edit.new_string, + ) + } else { + edit.new_string.clone() + }; + replacements.push(PlannedReplacement { + start, + end, + replacement, + edit_index, + }); + } + } + + replacements.sort_by_key(|replacement| (replacement.start, replacement.end)); + for pair in replacements.windows(2) { + if pair[1].start < pair[0].end { + return Err(format!( + "Edits {} and {} overlap in '{}'", + pair[0].edit_index + 1, + pair[1].edit_index + 1, + self.display_path(&path) + )); + } + } + + let mut updated = original.clone(); + for replacement in replacements.iter().rev() { + updated.replace_range(replacement.start..replacement.end, &replacement.replacement); + } + if updated == original { + return Err(format!( + "Edits would not change '{}'", + self.display_path(&path) + )); + } + + let display_path = self.display_path(&path); + let first_changed_line = first_changed_line(&original, &updated); + let (diff, patch) = build_diffs(&display_path, &original, &updated); + let restored = restore_document(&updated, line_ending, has_bom); + self.overlay.insert(path, OverlayEntry::File(restored)); + + Ok(EditOutcome { + display_path, + replacement_count: replacements.len(), + diff, + patch, + first_changed_line, + }) + } +} + +fn normalize_edit_text(text: &str) -> String { + text.strip_prefix('\u{feff}') + .unwrap_or(text) + .replace("\r\n", "\n") +} + +fn validate_edit(edit: &TextEdit, edit_index: usize) -> Result<(), String> { + if edit.old_string.is_empty() { + return Err(format!( + "Edit {} old_string must not be empty", + edit_index + 1 + )); + } + if edit.old_string == edit.new_string { + return Err(format!( + "Edit {} is a no-op because old_string and new_string are identical", + edit_index + 1 + )); + } + Ok(()) +} + +fn exact_ranges(content: &str, needle: &str) -> Vec<(usize, usize)> { + content + .match_indices(needle) + .map(|(start, _)| (start, start + needle.len())) + .collect() +} + +fn fuzzy_ranges(content: &str, needle: &str) -> Vec<(usize, usize)> { + let content_lines = line_spans(content); + let needle_has_trailing_newline = needle.ends_with('\n'); + let mut needle_lines = needle.split('\n').collect::>(); + if needle_has_trailing_newline { + needle_lines.pop(); + } + if needle_lines.is_empty() || needle_lines.len() > content_lines.len() { + return Vec::new(); + } + let normalized_needle = needle_lines + .iter() + .map(|line| normalize_fuzzy_line(line)) + .collect::>(); + + content_lines + .windows(normalized_needle.len()) + .enumerate() + .filter_map(|(index, window)| { + let equal = window + .iter() + .zip(&normalized_needle) + .all(|(line, needle)| normalize_fuzzy_line(line.text) == *needle); + if !equal { + return None; + } + let last = window.last()?; + if needle_has_trailing_newline && last.full_end == last.content_end { + return None; + } + Some(( + content_lines[index].start, + if needle_has_trailing_newline { + last.full_end + } else { + last.content_end + }, + )) + }) + .collect() +} + +fn line_spans(content: &str) -> Vec> { + if content.is_empty() { + return vec![LineSpan { + start: 0, + content_end: 0, + full_end: 0, + text: "", + }]; + } + + let mut spans = Vec::new(); + let mut start = 0; + for segment in content.split_inclusive('\n') { + let full_end = start + segment.len(); + let text = segment.strip_suffix('\n').unwrap_or(segment); + let content_end = start + text.len(); + spans.push(LineSpan { + start, + content_end, + full_end, + text, + }); + start = full_end; + } + if start < content.len() { + let text = &content[start..]; + spans.push(LineSpan { + start, + content_end: content.len(), + full_end: content.len(), + text, + }); + } + spans +} + +fn normalize_fuzzy_line(line: &str) -> String { + line.nfkc() + .map(|character| match character { + '\u{2018}' | '\u{2019}' | '\u{201A}' | '\u{201B}' => '\'', + '\u{201C}' | '\u{201D}' | '\u{201E}' | '\u{201F}' => '"', + '\u{2010}' | '\u{2011}' | '\u{2012}' | '\u{2013}' | '\u{2014}' | '\u{2015}' + | '\u{2212}' => '-', + other => other, + }) + .collect::() + .trim_end() + .to_string() +} + +fn overlay_fuzzy_replacement(matched: &str, old: &str, new: &str) -> String { + let matched_lines = matched.split('\n').collect::>(); + let old_lines = old.split('\n').collect::>(); + let new_lines = new.split('\n').collect::>(); + if matched_lines.len() != old_lines.len() || old_lines.len() != new_lines.len() { + return new.to_string(); + } + + matched_lines + .iter() + .zip(old_lines.iter().zip(new_lines.iter())) + .map(|(matched_line, (old_line, new_line))| { + if normalize_fuzzy_line(old_line) == normalize_fuzzy_line(new_line) { + (*matched_line).to_string() + } else { + (*new_line).to_string() + } + }) + .collect::>() + .join("\n") +} + +fn restore_document(content: &str, line_ending: LineEnding, has_bom: bool) -> Vec { + let restored = match line_ending { + LineEnding::Lf => content.to_string(), + LineEnding::CrLf => content.replace('\n', "\r\n"), + }; + let mut bytes = Vec::with_capacity(restored.len() + usize::from(has_bom) * UTF8_BOM.len()); + if has_bom { + bytes.extend_from_slice(UTF8_BOM); + } + bytes.extend_from_slice(restored.as_bytes()); + bytes +} + +fn first_changed_line(before: &str, after: &str) -> usize { + let before = before.split('\n').collect::>(); + let after = after.split('\n').collect::>(); + let shared = before.len().min(after.len()); + before + .iter() + .zip(&after) + .position(|(left, right)| left != right) + .map_or(shared + 1, |index| index + 1) +} + +fn build_diffs(path: &str, before: &str, after: &str) -> (String, String) { + let diff = TextDiff::from_lines(before, after); + let mut display = String::new(); + for change in diff.iter_all_changes() { + let marker = match change.tag() { + ChangeTag::Delete => '-', + ChangeTag::Insert => '+', + ChangeTag::Equal => ' ', + }; + display.push(marker); + display.push_str(change.value()); + if change.missing_newline() { + display.push('\n'); + } + } + let patch = diff.unified_diff().header(path, path).to_string(); + (display, patch) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn fuzzy_normalization_handles_nfkc_quotes_dashes_and_trailing_space() { + assert_eq!(normalize_fuzzy_line("A ‘quote’ — "), "A 'quote' -"); + } + + #[test] + fn fuzzy_overlay_preserves_unchanged_original_lines() { + let matched = "let label = “hello”; \nlet count = 1;"; + let old = "let label = \"hello\";\nlet count = 1;"; + let new = "let label = \"hello\";\nlet count = 2;"; + + assert_eq!( + overlay_fuzzy_replacement(matched, old, new), + "let label = “hello”; \nlet count = 2;" + ); + } +} diff --git a/vendor/mentra/src/tool/files/workspace/operations.rs b/vendor/mentra/src/tool/files/workspace/operations.rs new file mode 100644 index 0000000..dcca9d0 --- /dev/null +++ b/vendor/mentra/src/tool/files/workspace/operations.rs @@ -0,0 +1,423 @@ +use regex::RegexBuilder; + +use super::{BTreeSet, EntryKind, OverlayEntry, SearchOptions, WorkspaceEditor}; +use crate::tool::files::schema::{FileOperation, InsertPosition}; + +impl WorkspaceEditor { + pub(crate) fn apply_operation(&mut self, operation: FileOperation) -> Result { + match operation { + FileOperation::Read { + path, + offset, + limit, + } => self.read(path, offset.unwrap_or(1), limit.unwrap_or(2000)), + FileOperation::List { path, depth, limit } => { + self.list(path, depth.unwrap_or(1), limit.unwrap_or(200)) + } + FileOperation::Search { + path, + pattern, + limit, + } => self.search(path, &pattern, limit.unwrap_or(200)), + FileOperation::Create { path, content } => self.create(path, content), + FileOperation::Set { path, content } => self.set(path, content), + FileOperation::Replace { + path, + old, + new, + replace_all, + expected_replacements, + } => self.replace( + path, + &old, + &new, + replace_all.unwrap_or(false), + expected_replacements.unwrap_or(1), + ), + FileOperation::Insert { + path, + anchor, + position, + content, + occurrence, + } => self.insert(path, &anchor, position, &content, occurrence), + FileOperation::Move { from, to } => self.move_path(from, to), + FileOperation::Delete { path } => self.delete(path), + } + } + + pub(crate) fn read( + &mut self, + path: String, + offset: usize, + limit: usize, + ) -> Result { + if offset == 0 { + return Err("read offset must be at least 1".to_string()); + } + + let path = self.resolve_path(&path)?; + let path = self.authorize_read(&path, "files_read")?; + let content = self.load_text_file(&path)?; + let lines = content.lines().collect::>(); + let start = offset.saturating_sub(1).min(lines.len()); + let end = start.saturating_add(limit).min(lines.len()); + let numbered = lines[start..end] + .iter() + .enumerate() + .map(|(index, line)| format!("L{}: {}", start + index + 1, line)) + .collect::>(); + + let body = if numbered.is_empty() { + "(no lines)".to_string() + } else { + numbered.join("\n") + }; + Ok(format!("read {}\n{}", self.display_path(&path), body)) + } + + pub(crate) fn list(&self, path: String, depth: usize, limit: usize) -> Result { + let path = self.resolve_path(&path)?; + let path = self.authorize_read(&path, "files_list")?; + let kind = self.entry_kind(&path)?; + if kind == EntryKind::Missing { + return Err(format!( + "Path '{}' does not exist", + self.display_path(&path) + )); + } + + let mut entries = Vec::new(); + match kind { + EntryKind::File => { + entries.push(format!("[file] {}", self.display_relative_to(&path, &path))); + } + EntryKind::Dir => { + let mut visited = BTreeSet::new(); + if limit > 0 { + self.walk_entries( + &path, + 1, + depth, + "files_list", + &mut visited, + &mut |child, kind| { + let label = match kind { + EntryKind::File => "file", + EntryKind::Dir => "dir", + EntryKind::Missing => return Ok(true), + }; + entries.push(format!( + "[{label}] {}", + self.display_relative_to(&path, child) + )); + Ok(entries.len() < limit) + }, + )?; + } + } + EntryKind::Missing => unreachable!(), + } + + let body = if entries.is_empty() { + "(no entries)".to_string() + } else { + entries.join("\n") + }; + Ok(format!("list {}\n{}", self.display_path(&path), body)) + } + + pub(crate) fn search( + &self, + path: String, + pattern: &str, + limit: usize, + ) -> Result { + self.grep(path, pattern, SearchOptions::default(), limit) + } + + pub(crate) fn grep( + &self, + path: String, + pattern: &str, + options: SearchOptions, + limit: usize, + ) -> Result { + let path = self.resolve_path(&path)?; + let path = self.authorize_read(&path, "files_search")?; + let expression = if options.literal { + regex::escape(pattern) + } else { + pattern.to_string() + }; + let regex = RegexBuilder::new(&expression) + .case_insensitive(options.ignore_case) + .multi_line(options.multiline) + .dot_matches_new_line(options.multiline) + .build() + .map_err(|error| format!("Invalid regex pattern: {error}"))?; + let kind = self.entry_kind(&path)?; + if kind == EntryKind::Missing { + return Err(format!( + "Path '{}' does not exist", + self.display_path(&path) + )); + } + + let mut matches = Vec::new(); + match kind { + EntryKind::File => { + if self.path_matches_glob(&path, &path, options.file_glob.as_deref()) { + self.search_file(&path, ®ex, &options, limit, &mut matches)?; + } + } + EntryKind::Dir => { + self.collect_search_matches(&path, ®ex, &options, limit, &mut matches)? + } + EntryKind::Missing => unreachable!(), + } + + let body = if matches.is_empty() { + "(no matches)".to_string() + } else { + matches.join("\n") + }; + Ok(format!( + "search {} /{pattern}/\n{}", + self.display_path(&path), + body + )) + } + + pub(crate) fn glob(&self, path: String, pattern: &str, limit: usize) -> Result { + let path = self.resolve_path(&path)?; + let path = self.authorize_read(&path, "files_glob")?; + let kind = self.entry_kind(&path)?; + if kind == EntryKind::Missing { + return Err(format!( + "Path '{}' does not exist", + self.display_path(&path) + )); + } + + let mut matches = Vec::new(); + if limit > 0 { + match kind { + EntryKind::File => { + if self.path_matches_glob(&path, &path, Some(pattern)) { + matches.push(self.display_path(&path)); + } + } + EntryKind::Dir => { + self.collect_glob_matches(&path, pattern, limit, &mut matches)?; + } + EntryKind::Missing => unreachable!(), + } + } + + let body = if matches.is_empty() { + "(no matches)".to_string() + } else { + matches.join("\n") + }; + Ok(format!( + "glob {} /{pattern}/\n{}", + self.display_path(&path), + body + )) + } + + pub(crate) fn create(&mut self, path: String, content: String) -> Result { + let path = self.resolve_path(&path)?; + let path = self.authorize_write(&path, "files_write")?; + match self.entry_kind(&path)? { + EntryKind::Missing => { + self.overlay + .insert(path.clone(), OverlayEntry::File(content.into_bytes())); + Ok(format!("create {}", self.display_path(&path))) + } + EntryKind::File | EntryKind::Dir => Err(format!( + "Path '{}' already exists", + self.display_path(&path) + )), + } + } + + pub(crate) fn set(&mut self, path: String, content: String) -> Result { + let path = self.resolve_path(&path)?; + let path = self.authorize_write(&path, "files_write")?; + if self.entry_kind(&path)? != EntryKind::File { + return Err(format!( + "Path '{}' does not exist as a file", + self.display_path(&path) + )); + } + + self.overlay + .insert(path.clone(), OverlayEntry::File(content.into_bytes())); + Ok(format!("set {}", self.display_path(&path))) + } + + pub(crate) fn write(&mut self, path: String, content: String) -> Result { + let path = self.resolve_path(&path)?; + let path = self.authorize_write(&path, "files_write")?; + match self.entry_kind(&path)? { + EntryKind::Missing | EntryKind::File => { + let byte_count = content.len(); + self.overlay + .insert(path.clone(), OverlayEntry::File(content.into_bytes())); + Ok(format!( + "Wrote {byte_count} byte(s) to {}", + self.display_path(&path) + )) + } + EntryKind::Dir => Err(format!( + "Path '{}' is a directory", + self.display_path(&path) + )), + } + } + + pub(crate) fn replace( + &mut self, + path: String, + old: &str, + new: &str, + replace_all: bool, + expected_replacements: usize, + ) -> Result { + if old.is_empty() { + return Err("replace old text must not be empty".to_string()); + } + let path = self.resolve_path(&path)?; + let path = self.authorize_write(&path, "files_write")?; + let content = self.load_text_file(&path)?; + let actual_replacements = content.match_indices(old).count(); + if actual_replacements != expected_replacements { + return Err(format!( + "Expected {expected_replacements} replacement(s) in '{}', found {actual_replacements}", + self.display_path(&path) + )); + } + + let updated = if replace_all { + content.replace(old, new) + } else { + content.replacen(old, new, 1) + }; + + self.overlay + .insert(path.clone(), OverlayEntry::File(updated.into_bytes())); + Ok(format!( + "replace {} ({actual_replacements} replacement{})", + self.display_path(&path), + if actual_replacements == 1 { "" } else { "s" } + )) + } + + pub(crate) fn insert( + &mut self, + path: String, + anchor: &str, + position: InsertPosition, + content: &str, + occurrence: Option, + ) -> Result { + if anchor.is_empty() { + return Err("insert anchor must not be empty".to_string()); + } + + let path = self.resolve_path(&path)?; + let path = self.authorize_write(&path, "files_write")?; + let current = self.load_text_file(&path)?; + let locations = current + .match_indices(anchor) + .map(|(index, _)| index) + .collect::>(); + if locations.is_empty() { + return Err(format!( + "Anchor '{anchor}' was not found in '{}'", + self.display_path(&path) + )); + } + + let insert_at = match occurrence { + Some(occurrence) => { + if occurrence == 0 { + return Err("insert occurrence must be at least 1".to_string()); + } + locations.get(occurrence - 1).copied().ok_or_else(|| { + format!( + "Anchor occurrence {occurrence} was not found in '{}'", + self.display_path(&path) + ) + })? + } + None if locations.len() == 1 => locations[0], + None => { + return Err(format!( + "Anchor '{anchor}' is ambiguous in '{}' ({})", + self.display_path(&path), + locations.len() + )); + } + }; + + let insert_at = match position { + InsertPosition::Before => insert_at, + InsertPosition::After => insert_at + anchor.len(), + }; + let updated = format!( + "{}{}{}", + ¤t[..insert_at], + content, + ¤t[insert_at..] + ); + self.overlay + .insert(path.clone(), OverlayEntry::File(updated.into_bytes())); + Ok(format!("insert {}", self.display_path(&path))) + } + + pub(crate) fn move_path(&mut self, from: String, to: String) -> Result { + let from = self.resolve_path(&from)?; + let to = self.resolve_path(&to)?; + let from = self.authorize_write(&from, "files_write")?; + let to = self.authorize_write(&to, "files_write")?; + + if self.entry_kind(&from)? != EntryKind::File { + return Err(format!( + "Source '{}' does not exist as a file", + self.display_path(&from) + )); + } + if self.entry_kind(&to)? != EntryKind::Missing { + return Err(format!( + "Destination '{}' already exists", + self.display_path(&to) + )); + } + + let bytes = self.load_file_bytes(&from)?; + self.overlay.insert(from.clone(), OverlayEntry::Deleted); + self.overlay.insert(to.clone(), OverlayEntry::File(bytes)); + Ok(format!( + "move {} -> {}", + self.display_path(&from), + self.display_path(&to) + )) + } + + pub(crate) fn delete(&mut self, path: String) -> Result { + let path = self.resolve_path(&path)?; + let path = self.authorize_write(&path, "files_write")?; + if self.entry_kind(&path)? != EntryKind::File { + return Err(format!( + "Path '{}' does not exist as a file", + self.display_path(&path) + )); + } + + self.overlay.insert(path.clone(), OverlayEntry::Deleted); + Ok(format!("delete {}", self.display_path(&path))) + } +} diff --git a/vendor/mentra/src/tool/internal.rs b/vendor/mentra/src/tool/internal.rs new file mode 100644 index 0000000..b7d52c3 --- /dev/null +++ b/vendor/mentra/src/tool/internal.rs @@ -0,0 +1,46 @@ +use serde_json::Value; + +use crate::ContentBlock; + +use super::{ + RuntimeToolDescriptor, ToolApprovalCategory, ToolCapability, ToolDurability, + ToolExecutionCategory, ToolResult, ToolSideEffectLevel, +}; + +pub(crate) struct RuntimeDescriptorParts { + pub(crate) name: String, + pub(crate) description: String, + pub(crate) input_schema: Value, + pub(crate) capabilities: Vec, + pub(crate) side_effect_level: ToolSideEffectLevel, + pub(crate) durability: ToolDurability, + pub(crate) execution_category: ToolExecutionCategory, + pub(crate) approval_category: ToolApprovalCategory, +} + +pub(crate) fn build_runtime_descriptor(parts: RuntimeDescriptorParts) -> RuntimeToolDescriptor { + RuntimeToolDescriptor::builder(parts.name) + .description(parts.description) + .input_schema(parts.input_schema) + .capabilities(parts.capabilities) + .side_effect_level(parts.side_effect_level) + .durability(parts.durability) + .execution_category(parts.execution_category) + .approval_category(parts.approval_category) + .build() +} + +pub(crate) fn content_block_to_tool_result(surface: &str, block: ContentBlock) -> ToolResult { + match block { + ContentBlock::ToolResult { + content, is_error, .. + } => { + if is_error { + Err(content.to_display_string()) + } else { + Ok(content.to_display_string()) + } + } + _ => Err(format!("{surface} returned an unexpected content block")), + } +} diff --git a/vendor/mentra/src/tool/model.rs b/vendor/mentra/src/tool/model.rs new file mode 100644 index 0000000..d07ee81 --- /dev/null +++ b/vendor/mentra/src/tool/model.rs @@ -0,0 +1,662 @@ +use std::{ + any::Any, + path::{Path, PathBuf}, + sync::Arc, +}; + +use async_trait::async_trait; +use mentra_provider::ToolResultContent; +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +use crate::agent::{ + CompactionDetails, CompactionTrigger, DisposableSubagentTemplate, SpawnedAgentStatus, + SpawnedAgentSummary, +}; +use crate::runtime::{RuntimeError, TaskIntrinsicTool, TaskItem}; +use crate::team::{TeamDispatch, TeamMemberSummary, TeamMessage, TeamProtocolRequestSummary}; +use crate::tool::ToolAuthorizationPreview; + +use super::descriptor::{RuntimeToolDescriptor, ToolExecutionMode}; + +#[allow(unused_imports)] +pub use mentra_provider::ToolLoadingPolicy; +pub type ToolSpec = RuntimeToolDescriptor; + +#[cfg(test)] +mod tests { + use crate::tool::{ProviderToolSpec, ToolLoadingPolicy}; + use serde_json::json; + + #[test] + fn tool_spec_builder_defaults_to_immediate_loading() { + let spec = ProviderToolSpec::builder("echo_tool") + .description("Echo a value.") + .input_schema(json!({ + "type": "object", + "properties": { + "value": { "type": "string" } + } + })) + .build(); + + assert_eq!(spec.loading_policy, ToolLoadingPolicy::Immediate); + } + + #[test] + fn tool_spec_builder_supports_deferred_loading() { + let spec = ProviderToolSpec::builder("echo_tool") + .defer_loading(true) + .build(); + + assert_eq!(spec.loading_policy, ToolLoadingPolicy::Deferred); + } + + #[test] + fn tool_spec_deserialization_defaults_loading_policy() { + let spec: ProviderToolSpec = serde_json::from_value(json!({ + "name": "echo_tool", + "description": "Echo a value.", + "input_schema": { + "type": "object", + "properties": {} + } + })) + .expect("deserialize tool spec"); + + assert_eq!(spec.loading_policy, ToolLoadingPolicy::Immediate); + } +} + +/// A concrete tool call emitted by a model. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ToolCall { + pub id: String, + pub name: String, + pub input: Value, +} + +/// Execution context made available to a running tool. +pub struct ToolContext<'a> { + pub agent_id: String, + pub tool_call_id: String, + pub tool_name: String, + pub(crate) working_directory: PathBuf, + pub(crate) runtime: crate::runtime::RuntimeHandle, + pub(crate) agent: &'a mut crate::agent::Agent, + pub(crate) event_tx: crate::agent::AgentEventBus, + /// The options of the `Agent::run` call this execution is a step of. + /// Reachable only as [`child_run_options`](Self::child_run_options), so a + /// tool can share the run's aggregate bounds with work it spawns but cannot + /// read or edit the run's own policy. + pub(crate) run_options: crate::runtime::RunOptions, +} + +impl ToolContext<'_> { + pub fn working_directory(&self) -> &Path { + self.working_directory.as_path() + } + + /// [`RunOptions`](crate::runtime::RunOptions) for a run this tool spawns, + /// derived from the options of the run this tool is executing under — see + /// [`RunOptions::child`](crate::runtime::RunOptions::child) for what a child + /// inherits and what it resets. + /// + /// Thread these into the spawned run's own `Agent::run` call. A subagent + /// driven on `RunOptions::default()` instead gets a fresh, unbounded token + /// counter, so its spend escapes the parent's `token_budget` and a parent + /// cancel, stop, or deadline never reaches it. + pub fn child_run_options(&self) -> crate::runtime::RunOptions { + self.run_options.child() + } + + /// Emit a progress event for the currently executing tool. + pub fn emit_progress(&self, progress: String) { + self.event_tx + .send(crate::agent::AgentEvent::ToolExecutionProgress { + id: self.tool_call_id.clone(), + name: self.tool_name.clone(), + progress, + }); + } + + pub fn agent_name(&self) -> &str { + self.agent.name() + } + + pub fn model(&self) -> &str { + self.agent.model() + } + + pub fn history_len(&self) -> usize { + self.agent.history().len() + } + + pub fn tasks(&self) -> &[TaskItem] { + self.agent.tasks() + } + + /// Returns this agent's tool-result paging configuration, if enabled. + pub(crate) fn tool_result_paging(&self) -> Option { + self.agent.config().tool_result_paging + } + + /// Returns the full text of one of this agent's paged tool results. + /// Scoped to the agent by construction: the retained results live on the + /// agent this context borrows, so no cross-agent read is expressible. + pub(crate) fn paged_tool_result(&self, tool_use_id: &str) -> Option> { + self.agent.paged_tool_result(tool_use_id) + } + + pub fn resolve_working_directory( + &self, + working_directory: Option<&str>, + ) -> Result { + self.runtime + .resolve_working_directory(&self.agent_id, working_directory) + } + + pub fn load_skill(&self, name: &str) -> Result { + self.runtime.load_skill(name) + } + + pub fn skill_descriptions(&self) -> Option { + self.runtime.skill_descriptions() + } + + pub fn app_context(&self) -> Result, String> + where + T: Any + Send + Sync + 'static, + { + self.runtime.app_context::() + } + + /// Runs one command on the local executor. + pub async fn execute_shell_command( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + self.runtime + .execute_shell_command( + &self.agent_id, + command, + justification, + requested_timeout, + cwd, + ) + .await + } + + /// Runs one command on the executor the host named. + /// + /// A tool that lets its caller say *where* a command runs passes the name + /// here; `None` is the local executor. The name reaches the installed + /// [`crate::runtime::RuntimeExecutor`] on the request and is interpreted + /// only there, so a tool can route a command without gaining any way to + /// route around the policy that authorized it. + pub async fn execute_shell_command_on( + &self, + target: Option, + command: String, + justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + self.runtime + .execute_shell_command_on( + &self.agent_id, + target, + command, + justification, + requested_timeout, + cwd, + ) + .await + } + + pub fn start_background_task( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + self.runtime.start_background_task( + &self.agent_id, + command, + justification, + requested_timeout, + cwd, + ) + } + + pub fn check_background_task(&self, task_id: Option<&str>) -> Result { + self.runtime.check_background_task(&self.agent_id, task_id) + } + + pub fn request_idle(&mut self) { + self.agent.request_idle(); + } + + pub async fn compact_history(&mut self) -> Result, RuntimeError> { + self.agent + .compact_history( + self.agent.history().len().saturating_sub(1), + CompactionTrigger::Manual, + ) + .await + } + + pub fn execute_task_tool( + &self, + tool: &TaskIntrinsicTool, + input: Value, + ) -> Result { + self.agent.execute_task_mutation(tool, input) + } + + pub fn refresh_tasks(&mut self) -> Result<(), RuntimeError> { + self.agent.refresh_tasks_from_disk() + } + + pub async fn read_file(&self, path: &str, max_lines: Option) -> Result { + self.runtime + .read_file(&self.agent_id, path, max_lines) + .await + } + + pub fn spawn_subagent(&self) -> Result { + self.agent.spawn_subagent() + } + + pub fn register_subagent(&mut self, agent: &crate::agent::Agent) -> SpawnedAgentSummary { + self.agent.register_subagent(agent) + } + + pub fn finish_subagent( + &mut self, + id: &str, + status: SpawnedAgentStatus, + ) -> Option { + self.agent.finish_subagent(id, status) + } + + pub async fn spawn_teammate( + &mut self, + name: impl Into, + role: impl Into, + prompt: Option, + ) -> Result { + self.agent.spawn_teammate(name, role, prompt).await + } + + pub fn send_team_message( + &self, + to: &str, + content: impl Into, + ) -> Result { + self.agent.send_team_message(to, content) + } + + pub fn broadcast_team_message( + &self, + content: impl Into, + ) -> Result, RuntimeError> { + self.agent.broadcast_team_message(content) + } + + pub fn read_team_inbox(&self) -> Result, RuntimeError> { + self.agent.read_team_inbox() + } + + pub fn request_team_protocol( + &self, + to: &str, + protocol: impl Into, + content: impl Into, + ) -> Result { + self.agent.request_team_protocol(to, protocol, content) + } + + pub fn respond_team_protocol( + &self, + request_id: &str, + approve: bool, + reason: Option, + ) -> Result { + self.agent + .respond_team_protocol(request_id, approve, reason) + } +} + +/// Execution context made available to a parallel-safe running tool. +#[derive(Clone)] +pub struct ParallelToolContext { + pub agent_id: String, + pub tool_call_id: String, + pub tool_name: String, + pub(crate) working_directory: PathBuf, + pub(crate) runtime: crate::runtime::RuntimeHandle, + pub(crate) subagent_template: DisposableSubagentTemplate, + pub(crate) agent_name: String, + pub(crate) model: String, + pub(crate) history_len: usize, + pub(crate) tasks: Vec, + pub(crate) event_tx: crate::agent::AgentEventBus, + /// The options of the `Agent::run` call this execution is a step of, as on + /// [`ToolContext::run_options`]. + pub(crate) run_options: crate::runtime::RunOptions, +} + +impl From> for ParallelToolContext { + fn from(ctx: ToolContext) -> Self { + ParallelToolContext { + agent_id: ctx.agent_id, + tool_call_id: ctx.tool_call_id, + tool_name: ctx.tool_name, + working_directory: ctx.working_directory, + runtime: ctx.runtime, + subagent_template: ctx.agent.disposable_subagent_template(), + agent_name: ctx.agent.name().to_string(), + model: ctx.agent.model().to_string(), + history_len: ctx.agent.history().len(), + tasks: ctx.agent.tasks().to_vec(), + event_tx: ctx.event_tx, + run_options: ctx.run_options, + } + } +} + +impl ParallelToolContext { + pub fn working_directory(&self) -> &Path { + self.working_directory.as_path() + } + + /// Emit a progress event for the currently executing tool. + pub fn emit_progress(&self, progress: String) { + self.event_tx + .send(crate::agent::AgentEvent::ToolExecutionProgress { + id: self.tool_call_id.clone(), + name: self.tool_name.clone(), + progress, + }); + } + + pub fn agent_name(&self) -> &str { + &self.agent_name + } + + pub fn model(&self) -> &str { + &self.model + } + + pub fn history_len(&self) -> usize { + self.history_len + } + + pub fn tasks(&self) -> &[TaskItem] { + &self.tasks + } + + pub fn resolve_working_directory( + &self, + working_directory: Option<&str>, + ) -> Result { + self.runtime + .resolve_working_directory(&self.agent_id, working_directory) + } + + pub(crate) fn shell_validation( + &self, + command: &str, + ) -> Result { + self.runtime.shell_validation(&self.agent_id, command) + } + + pub fn load_skill(&self, name: &str) -> Result { + self.runtime.load_skill(name) + } + + pub fn skill_descriptions(&self) -> Option { + self.runtime.skill_descriptions() + } + + pub fn app_context(&self) -> Result, String> + where + T: Any + Send + Sync + 'static, + { + self.runtime.app_context::() + } + + /// Runs one command on the local executor. + pub async fn execute_shell_command( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + self.runtime + .execute_shell_command( + &self.agent_id, + command, + justification, + requested_timeout, + cwd, + ) + .await + } + + /// Runs one command on the executor the host named. + /// + /// A tool that lets its caller say *where* a command runs passes the name + /// here; `None` is the local executor. The name reaches the installed + /// [`crate::runtime::RuntimeExecutor`] on the request and is interpreted + /// only there, so a tool can route a command without gaining any way to + /// route around the policy that authorized it. + pub async fn execute_shell_command_on( + &self, + target: Option, + command: String, + justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + self.runtime + .execute_shell_command_on( + &self.agent_id, + target, + command, + justification, + requested_timeout, + cwd, + ) + .await + } + + pub fn start_background_task( + &self, + command: String, + justification: Option, + requested_timeout: Option, + cwd: PathBuf, + ) -> Result { + self.runtime.start_background_task( + &self.agent_id, + command, + justification, + requested_timeout, + cwd, + ) + } + + pub fn check_background_task(&self, task_id: Option<&str>) -> Result { + self.runtime.check_background_task(&self.agent_id, task_id) + } + + pub async fn read_file(&self, path: &str, max_lines: Option) -> Result { + self.runtime + .read_file(&self.agent_id, path, max_lines) + .await + } + + pub fn spawn_subagent(&self) -> Result { + self.subagent_template.spawn() + } + + /// [`RunOptions`](crate::runtime::RunOptions) for a run this tool spawns — + /// the parallel-lane counterpart of + /// [`ToolContext::child_run_options`], carrying the same caveat: a subagent + /// from [`spawn_subagent`](Self::spawn_subagent) driven on + /// `RunOptions::default()` gets a fresh, unbounded token counter and shares + /// none of the parent run's cancellation, stop, or deadline. + pub fn child_run_options(&self) -> crate::runtime::RunOptions { + self.run_options.child() + } +} + +/// String result returned by Mentra tools. +pub type ToolResult = Result; + +/// Structured, additive successor to [`ToolResult`]. +/// +/// `content` is the provider-visible projection of a tool's result and reuses +/// the existing [`ToolResultContent`] from `mentra-provider`, so no new +/// provider representation is required. `details` is opaque host metadata +/// that survives the local transcript but is never sent to a provider — mentra +/// never interprets it. `terminate` asks the run to end as the value of this +/// tool's own execution: a first-class successor to +/// [`ToolContext::request_idle`] for terminal actions, honored only when the +/// call executes in an exclusive lane (see [`RuntimeToolDescriptorBuilder::terminal`]). +/// +/// Tool-level failures keep using the existing `Err(String)` channel on +/// [`ToolExecutor::execute_output`] / [`ToolExecutor::execute_mut_output`]; +/// `ToolOutput` only ever appears on the `Ok` side. +#[derive(Debug, Clone)] +pub struct ToolOutput { + pub content: ToolResultContent, + pub details: Option, + pub terminate: bool, +} + +impl ToolOutput { + /// Builds a plain-text, non-terminating output with no attached metadata. + pub fn text(content: impl Into) -> Self { + Self { + content: ToolResultContent::Text(content.into()), + details: None, + terminate: false, + } + } + + /// Builds a structured, non-terminating output with no attached metadata. + pub fn structured(content: Value) -> Self { + Self { + content: ToolResultContent::Structured(content), + details: None, + terminate: false, + } + } + + /// Attaches opaque host metadata that survives transcript persistence but + /// is never projected to a provider. + pub fn with_details(mut self, details: Value) -> Self { + self.details = Some(details); + self + } + + /// Marks this output as ending the run as the value of its own execution. + pub fn terminating(mut self) -> Self { + self.terminate = true; + self + } +} + +/// Bridges an existing `Ok(String)` tool result into the additive structured +/// path: `Text` content, no metadata, no termination. +impl From for ToolOutput { + fn from(value: String) -> Self { + Self::text(value) + } +} + +/// Definition contract for custom tools exposed to models. +pub trait ToolDefinition: Send + Sync { + fn descriptor(&self) -> RuntimeToolDescriptor; +} + +/// Execution contract for custom tools exposed to models. +#[async_trait] +pub trait ToolExecutor: ToolDefinition + Send + Sync { + fn authorization_preview( + &self, + ctx: &ParallelToolContext, + input: &Value, + ) -> Result { + let descriptor = self.descriptor(); + Ok(ToolAuthorizationPreview { + working_directory: ctx.working_directory().to_path_buf(), + capabilities: descriptor.capabilities, + side_effect_level: descriptor.side_effect_level, + durability: descriptor.durability, + execution_category: descriptor.execution_category, + approval_category: descriptor.approval_category, + raw_input: input.clone(), + structured_input: input.clone(), + }) + } + + fn execution_category(&self, _input: &Value) -> super::descriptor::ToolExecutionCategory { + self.descriptor().execution_category + } + + fn execution_mode(&self, input: &Value) -> ToolExecutionMode { + self.execution_category(input).into() + } + + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + Err(format!( + "Tool '{}' does not support parallel execution", + self.descriptor().provider.name + )) + } + + async fn execute_mut(&self, ctx: ToolContext<'_>, input: Value) -> ToolResult { + self.execute(ctx.into(), input).await + } + + /// Structured, parallel-lane execution. Defaults to bridging + /// [`ToolExecutor::execute`] through `ToolOutput::from`, so every + /// existing string-returning tool keeps working unchanged. Overriding + /// this directly (instead of `execute`) opts a tool into structured + /// content, opaque details, or (subject to the exclusive-lane + /// requirement) termination. + async fn execute_output( + &self, + ctx: ParallelToolContext, + input: Value, + ) -> Result { + self.execute(ctx, input).await.map(ToolOutput::from) + } + + /// Structured, exclusive-lane execution. Defaults to bridging + /// [`ToolExecutor::execute_mut`] through `ToolOutput::from`, so every + /// existing string-returning tool keeps working unchanged. Overriding + /// this directly (instead of `execute_mut`) opts a tool into structured + /// content, opaque details, or termination. + async fn execute_mut_output( + &self, + ctx: ToolContext<'_>, + input: Value, + ) -> Result { + self.execute_mut(ctx, input).await.map(ToolOutput::from) + } +} + +/// Runtime tool contract used by Mentra registries and execution. +pub trait ExecutableTool: ToolDefinition + ToolExecutor {} + +impl ExecutableTool for T where T: ToolDefinition + ToolExecutor {} diff --git a/vendor/mentra/src/tool/orchestrator.rs b/vendor/mentra/src/tool/orchestrator.rs new file mode 100644 index 0000000..8c752e0 --- /dev/null +++ b/vendor/mentra/src/tool/orchestrator.rs @@ -0,0 +1,1059 @@ +//! Tool orchestration pipeline for scheduling, authorization, execution, and result ordering. + +use std::{collections::BTreeMap, future::Future, path::PathBuf, sync::Arc, time::Duration}; + +use tokio::task::JoinSet; + +use crate::{ + ContentBlock, + agent::{Agent, AgentEvent, AgentStatus}, + error::RuntimeError, + runtime::control::{HookDecision, PreExecutionContext}, + runtime::{RunOptions, RuntimeHookEvent}, + tool::{ + ExecutableTool, ParallelToolContext, RuntimeToolDescriptor, ToolAuthorizationOutcome, + ToolAuthorizationRequest, ToolCall, ToolCapability, ToolContext, ToolExecutionCategory, + }, +}; + +use super::{ + paging::{READ_TOOL_RESULT_TOOL, ToolResultPager}, + truncation::{SpillBehavior, ToolOutputLimiter}, +}; + +const PARALLEL_JOIN_POLL_INTERVAL: Duration = Duration::from_millis(10); + +pub(crate) struct ToolExecutionOutcome { + pub(crate) results: Vec, + pub(crate) successful_task: bool, + pub(crate) end_turn: bool, + /// Per-call opaque metadata collected from this round's executions, + /// keyed by `tool_use_id` — the runner attaches this to the appended + /// transcript item so it survives persistence and replay, never + /// projected to a provider (ADR-0001 §4). + pub(crate) details: BTreeMap, +} + +pub(crate) struct ToolRuntime { + runtime: crate::runtime::handle::RuntimeHandle, + agent_id: String, + tool_calls: usize, + working_directory: Option, + output_limiter: ToolOutputLimiter, + /// `Some` only when this agent enables tool-result paging; `None` leaves + /// every result exactly as the limiter produced it. + pager: Option, +} + +#[derive(Clone)] +enum ToolCallBatch { + Exclusive(ToolCall), + Parallel(Vec), +} + +struct ToolCallSchedule { + batches: Vec, +} + +struct CompletedToolExecution { + result: ContentBlock, + task_succeeded: bool, + /// Ends the current round: true when this execution consumed + /// [`crate::tool::ToolContext::request_idle`] (exclusive lane) or its + /// [`crate::tool::ToolOutput::terminate`] successor. Controls whether + /// `TurnRunner::run` issues another model round. + should_end_turn: bool, + /// True only when `should_end_turn` came from `ToolOutput::terminate` + /// specifically (never from the pre-existing idle-request signal). + /// Distinct from `should_end_turn` because it additionally drives + /// skipping not-yet-executed batches later in the same round — a new + /// behavior scoped to genuine termination, not to idle requests, so + /// existing `request_idle` callers see unchanged behavior. + terminated: bool, + tool_name: String, + /// This execution's opaque `ToolOutput::details`, if any — collected by + /// [`ToolRuntime::execute_calls`] into [`ToolExecutionOutcome::details`]. + details: Option, +} + +/// How a single execution affects the current round — bundled so +/// [`ToolRuntime::completed_execution`] stays within a reasonable argument +/// count. `Default` is "continues": neither ends the round nor terminates. +#[derive(Debug, Clone, Copy, Default)] +struct RoundEffect { + should_end_turn: bool, + terminated: bool, +} + +impl ToolRuntime { + pub(crate) fn new(agent: &Agent) -> Self { + let runtime = agent.runtime_handle(); + let policy = &runtime.execution.policy; + let spill = if !policy.spill_full_tool_output { + SpillBehavior::Disabled("spill-to-file is disabled by runtime policy") + } else if !runtime.persistence.store.allows_disk_artifacts() { + SpillBehavior::Disabled("the runtime store forbids durable artifacts") + } else { + SpillBehavior::Enabled(agent.config().compaction.transcript_dir.join("tool-output")) + }; + let output_limiter = ToolOutputLimiter::new( + policy.max_tool_result_bytes, + policy.max_tool_result_lines, + spill, + ); + Self { + runtime, + agent_id: agent.id().to_string(), + tool_calls: 0, + working_directory: None, + output_limiter, + pager: agent.config().tool_result_paging.map(ToolResultPager::new), + } + } + + pub(crate) async fn execute_calls( + &mut self, + agent: &mut Agent, + options: &RunOptions, + calls: Vec, + ) -> Result { + let mut results = Vec::new(); + let mut successful_task = false; + let mut end_turn = false; + let mut details = BTreeMap::new(); + + let mut batches = ToolCallSchedule::new(self, agent, calls) + .batches + .into_iter(); + + while let Some(batch) = batches.next() { + options.check_limits()?; + let execution_count = batch.execution_count(); + if self.tool_calls + execution_count > options.tool_budget() { + return Err(RuntimeError::ToolBudgetExceeded(options.tool_budget())); + } + self.tool_calls += execution_count; + + let executions = match batch { + ToolCallBatch::Exclusive(call) => { + vec![self.execute_one_tool(agent, options, call).await?] + } + ToolCallBatch::Parallel(calls) => { + self.execute_parallel_batch(agent, options, calls).await? + } + }; + + let mut terminator = None; + for execution in executions { + successful_task |= execution.task_succeeded; + end_turn |= execution.should_end_turn; + let result = self.page_result(agent, &execution.tool_name, execution.result); + if execution.terminated { + terminator.get_or_insert(execution.tool_name); + } + if let (Some(value), ContentBlock::ToolResult { tool_use_id, .. }) = + (execution.details, &result) + { + details.insert(tool_use_id.clone(), value); + } + results.push(result); + } + + // A terminating call ends the round as the value of its own + // execution; calls already scheduled for later batches in this + // round are never executed. Each still gets an explicit + // is_error result so the transcript always has one result block + // per tool_use — never a silent drop. + if let Some(terminator) = terminator { + for remaining_batch in batches { + for call in remaining_batch.into_calls() { + let result = not_executed_result(&call, &terminator); + results.push(self.page_result(agent, &call.name, result)); + } + } + break; + } + } + + Ok(ToolExecutionOutcome { + results, + successful_task, + end_turn, + details, + }) + } + + /// Replaces an oversized text result with its first window, retaining the + /// full text on the agent for `read_tool_result` to serve. + /// + /// This is the single point where a result becomes the *model's* view of + /// itself: every `AgentEvent::ToolExecutionFinished` has already been + /// emitted with the complete block by the time a result reaches here, so + /// consumers reconstructing evidence from the event stream observe no + /// change at all. Applied to every block that joins the round's committed + /// message — including the fixed not-executed and not-found results — so + /// no path into the transcript bypasses the bound. + fn page_result(&self, agent: &Agent, tool_name: &str, result: ContentBlock) -> ContentBlock { + let Some(pager) = self.pager else { + return result; + }; + // A window returned by `read_tool_result` is bounded by construction; + // paging it again would nest a trailer inside a trailer. + if tool_name == READ_TOOL_RESULT_TOOL { + return result; + } + let ContentBlock::ToolResult { + tool_use_id, + content: mentra_provider::ToolResultContent::Text(text), + is_error, + } = result + else { + return result; + }; + + let Some(page) = pager.first_page(&tool_use_id, &text) else { + return ContentBlock::ToolResult { + tool_use_id, + content: mentra_provider::ToolResultContent::Text(text), + is_error, + }; + }; + agent.record_paged_tool_result(&tool_use_id, &text); + ContentBlock::ToolResult { + tool_use_id, + content: mentra_provider::ToolResultContent::Text(page), + is_error, + } + } + + fn call_execution_category_for_agent( + &self, + call: &ToolCall, + agent: Option<&Agent>, + ) -> ToolExecutionCategory { + if agent.is_some_and(|agent| !agent.can_use_tool(&call.name)) { + return ToolExecutionCategory::ExclusiveLocalMutation; + } + + let Some(tool) = self.runtime.get_tool(&call.name) else { + return ToolExecutionCategory::ExclusiveLocalMutation; + }; + let category = tool.execution_category(&call.input); + let terminal = self + .runtime + .get_tool_descriptor(&call.name) + .is_some_and(|descriptor| descriptor.terminal); + + // STATIC exclusivity: a terminal-marked tool is never scheduled in a + // parallel batch, regardless of its declared execution_category — + // coerce rather than panic, matching the existing fallback-to-exclusive + // precedent above. + if terminal && category.allows_parallel() { + eprintln!( + "warning: tool '{}' is marked terminal but declared a parallel \ + execution category; coercing to exclusive scheduling", + call.name + ); + return ToolExecutionCategory::ExclusiveLocalMutation; + } + + category + } + + fn note_tool_started( + &mut self, + agent: &mut Agent, + call: &ToolCall, + ) -> Result<(), RuntimeError> { + agent.set_status(AgentStatus::ExecutingTool { + id: call.id.clone(), + name: call.name.clone(), + }); + agent.emit_event(AgentEvent::ToolExecutionStarted { call: call.clone() }); + agent.update_run_state("executing_tool", None) + } + + fn emit_tool_runtime_started(&self, call: &ToolCall) -> Result<(), RuntimeError> { + self.runtime + .emit_hook(RuntimeHookEvent::ToolExecutionStarted { + agent_id: self.agent_id.clone(), + tool_name: call.name.clone(), + tool_call_id: call.id.clone(), + }) + } + + fn emit_tool_runtime_finished( + &self, + call: &ToolCall, + result: &ContentBlock, + details: Option, + ) { + let is_error = matches!(result, ContentBlock::ToolResult { is_error: true, .. }); + let output_preview = match result { + ContentBlock::ToolResult { content, .. } => content.to_display_string(), + _ => String::new(), + }; + let error = is_error.then_some(output_preview.clone()); + let _ = self + .runtime + .emit_hook(RuntimeHookEvent::ToolExecutionFinished { + agent_id: self.agent_id.clone(), + tool_name: call.name.clone(), + tool_call_id: call.id.clone(), + is_error, + error, + output_preview, + details, + }); + } + + fn emit_tool_authorization_started( + &self, + call: &ToolCall, + preview: crate::tool::ToolAuthorizationPreview, + ) -> Result<(), RuntimeError> { + self.runtime + .emit_hook(RuntimeHookEvent::ToolAuthorizationStarted { + agent_id: self.agent_id.clone(), + tool_name: call.name.clone(), + tool_call_id: call.id.clone(), + preview, + }) + } + + fn emit_tool_authorization_finished( + &self, + call: &ToolCall, + outcome: ToolAuthorizationOutcome, + reason: Option, + ) -> Result<(), RuntimeError> { + self.runtime + .emit_hook(RuntimeHookEvent::ToolAuthorizationFinished { + agent_id: self.agent_id.clone(), + tool_name: call.name.clone(), + tool_call_id: call.id.clone(), + outcome, + reason, + }) + } + + fn emit_tool_authorization_blocked( + &self, + call: &ToolCall, + outcome: ToolAuthorizationOutcome, + reason: Option, + ) -> Result<(), RuntimeError> { + self.runtime + .emit_hook(RuntimeHookEvent::ToolAuthorizationBlocked { + agent_id: self.agent_id.clone(), + tool_name: call.name.clone(), + tool_call_id: call.id.clone(), + outcome, + reason, + }) + } + + async fn run_pre_hooks(&mut self, call: &ToolCall) -> Result { + let context = PreExecutionContext { + agent_id: self.agent_id.clone(), + tool_name: call.name.clone(), + tool_call_id: call.id.clone(), + input_json: serde_json::to_string(&call.input).unwrap_or_default(), + working_directory: self.working_directory(), + }; + self.runtime.pre_hooks().run(&context).await + } + + /// Runs the pre-execution hooks and applies whatever they decided. + /// + /// `Ok(None)` means proceed — possibly with `call.input` rewritten by a + /// hook. `Ok(Some(reason))` means the call must not run and `reason` is + /// what the model should be told. + /// + /// Shared by the serial and parallel paths so the two cannot disagree + /// about what a hook's answer means. + async fn apply_pre_hooks( + &mut self, + call: &mut ToolCall, + ) -> Result, RuntimeError> { + match self.run_pre_hooks(call).await? { + HookDecision::Allow => Ok(None), + HookDecision::Deny(reason) => Ok(Some(reason)), + HookDecision::Modify { input_json, .. } => { + // A hook that rewrites the input but hands back something that + // is not JSON has failed at its own job. Refusing is the safe + // reading: running the *original* would silently ignore a hook + // that believed it had intervened. + match serde_json::from_str(&input_json) { + Ok(input) => { + call.input = input; + Ok(None) + } + Err(error) => Ok(Some(format!( + "pre-execution hook returned invalid JSON for '{}': {error}", + call.name + ))), + } + } + } + } + + fn emit_tool_execution_blocked(&self, call: &ToolCall, reason: &str) { + let _ = self + .runtime + .emit_hook(RuntimeHookEvent::ToolExecutionBlocked { + agent_id: self.agent_id.clone(), + tool_name: call.name.clone(), + tool_call_id: call.id.clone(), + reason: reason.to_string(), + }); + } + + fn unavailable_tool_result(&self, call: ToolCall) -> ContentBlock { + ContentBlock::ToolResult { + tool_use_id: call.id, + content: format!("Tool '{}' is not available for this agent", call.name).into(), + is_error: true, + } + } + + fn blocked_tool_result(&self, call: &ToolCall, error: RuntimeError) -> ContentBlock { + ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: format!("Tool execution blocked: {error}").into(), + is_error: true, + } + } + + fn blocked_authorization_result( + &self, + call: &ToolCall, + outcome: ToolAuthorizationOutcome, + reason: Option, + ) -> ContentBlock { + let content = match outcome { + ToolAuthorizationOutcome::Allow => "Tool execution blocked by authorizer".to_string(), + ToolAuthorizationOutcome::Prompt => reason + .map(|reason| format!("Tool execution requires approval: {reason}")) + .unwrap_or_else(|| "Tool execution requires approval".to_string()), + ToolAuthorizationOutcome::Deny => reason + .map(|reason| format!("Tool execution denied: {reason}")) + .unwrap_or_else(|| "Tool execution denied by authorizer".to_string()), + }; + + ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: content.into(), + is_error: true, + } + } + + /// Splits a structured tool outcome into its provider-visible + /// projection, opaque host metadata, and requested termination — the + /// single boundary where `details` is separated from what a provider + /// ever sees (only `content` reaches `ContentBlock::ToolResult`). + async fn tool_output_block( + &self, + call: &ToolCall, + output: Result, + ) -> (ContentBlock, Option, bool) { + match output { + Ok(output) => ( + ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: self.output_limiter.apply(output.content).await, + is_error: false, + }, + output.details, + output.terminate, + ), + Err(content) => ( + ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: self + .output_limiter + .apply(mentra_provider::ToolResultContent::Text(content)) + .await, + is_error: true, + }, + None, + false, + ), + } + } + + fn completed_execution( + &self, + agent: &Agent, + call: &ToolCall, + descriptor: &RuntimeToolDescriptor, + result: ContentBlock, + effect: RoundEffect, + details: Option, + ) -> CompletedToolExecution { + self.emit_tool_runtime_finished(call, &result, details.clone()); + agent.emit_event(AgentEvent::ToolExecutionFinished { + result: result.clone(), + }); + let task_succeeded = matches!( + &result, + ContentBlock::ToolResult { + is_error: false, + .. + } + ) && descriptor + .capabilities + .iter() + .any(|capability| matches!(capability, ToolCapability::TaskMutation)); + + CompletedToolExecution { + result, + task_succeeded, + should_end_turn: effect.should_end_turn, + terminated: effect.terminated, + tool_name: call.name.clone(), + details, + } + } + + fn working_directory(&mut self) -> std::path::PathBuf { + if let Some(path) = &self.working_directory { + return path.clone(); + } + + let path = self + .runtime + .resolve_working_directory(&self.agent_id, None) + .unwrap_or_else(|_| self.runtime.default_working_directory(&self.agent_id)); + self.working_directory = Some(path.clone()); + path + } + + fn parallel_tool_context( + &mut self, + agent: &Agent, + options: &RunOptions, + call: &ToolCall, + ) -> ParallelToolContext { + ParallelToolContext { + agent_id: self.agent_id.clone(), + tool_call_id: call.id.clone(), + tool_name: call.name.clone(), + working_directory: self.working_directory(), + runtime: self.runtime.clone(), + subagent_template: agent.disposable_subagent_template(), + agent_name: agent.name().to_string(), + model: agent.model().to_string(), + history_len: agent.history().len(), + tasks: agent.tasks().to_vec(), + event_tx: agent.event_sender(), + run_options: options.clone(), + } + } + + fn registered_tool( + &self, + name: &str, + ) -> Option<(Arc, RuntimeToolDescriptor)> { + let tool = self.runtime.get_tool(name)?; + let descriptor = self.runtime.get_tool_descriptor(name)?; + Some((tool, descriptor)) + } + + async fn authorize_tool_call( + &self, + call: &ToolCall, + tool: &Arc, + ctx: &ParallelToolContext, + ) -> Result, RuntimeError> { + let Some(authorizer) = self.runtime.execution.tool_authorizer.clone() else { + return Ok(None); + }; + + let preview = match tool.authorization_preview(ctx, &call.input) { + Ok(preview) => preview, + Err(error) => { + return Ok(Some(self.blocked_authorization_result( + call, + ToolAuthorizationOutcome::Deny, + Some(error), + ))); + } + }; + + self.emit_tool_authorization_started(call, preview.clone())?; + let request = ToolAuthorizationRequest { + agent_id: self.agent_id.clone(), + agent_name: ctx.agent_name().to_string(), + model: ctx.model().to_string(), + history_len: ctx.history_len(), + tool_call_id: call.id.clone(), + tool_name: call.name.clone(), + preview, + }; + + let result = match authorizer.timeout() { + Some(timeout) => { + match tokio::time::timeout(timeout, authorizer.authorize(&request)).await { + Ok(result) => result, + Err(_) => { + return self.handle_authorization_block( + call, + ToolAuthorizationOutcome::Deny, + Some(format!( + "authorizer timed out after {}", + format_duration(timeout) + )), + ); + } + } + } + None => authorizer.authorize(&request).await, + }; + + match result { + Ok(decision) => match decision.outcome { + ToolAuthorizationOutcome::Allow => { + self.emit_tool_authorization_finished(call, decision.outcome, decision.reason)?; + Ok(None) + } + outcome => self.handle_authorization_block(call, outcome, decision.reason), + }, + Err(error) => self.handle_authorization_block( + call, + ToolAuthorizationOutcome::Deny, + Some(error.to_string()), + ), + } + } + + fn handle_authorization_block( + &self, + call: &ToolCall, + outcome: ToolAuthorizationOutcome, + reason: Option, + ) -> Result, RuntimeError> { + self.emit_tool_authorization_finished(call, outcome, reason.clone())?; + self.emit_tool_authorization_blocked(call, outcome, reason.clone())?; + Ok(Some( + self.blocked_authorization_result(call, outcome, reason), + )) + } + + async fn execute_one_tool( + &mut self, + agent: &mut Agent, + options: &RunOptions, + call: ToolCall, + ) -> Result { + self.note_tool_started(agent, &call)?; + if !agent.can_use_tool(&call.name) { + let result = self.unavailable_tool_result(call.clone()); + agent.emit_event(AgentEvent::ToolExecutionFinished { + result: result.clone(), + }); + return Ok(CompletedToolExecution { + result, + task_succeeded: false, + should_end_turn: false, + terminated: false, + tool_name: call.name.clone(), + details: None, + }); + } + + Ok(self.execute_registered_tool(agent, options, call).await) + } + + async fn execute_parallel_batch( + &mut self, + agent: &mut Agent, + options: &RunOptions, + calls: Vec, + ) -> Result, RuntimeError> { + let len = calls.len(); + let mut results = (0..len).map(|_| None).collect::>(); + let mut join_set = JoinSet::new(); + + for (index, mut call) in calls.iter().cloned().enumerate() { + if let Err(error) = self.note_tool_started(agent, &call) { + join_set.abort_all(); + return Err(error); + } + + let Some((tool, descriptor)) = self.registered_tool(&call.name) else { + let result = ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: "Tool not found".into(), + is_error: true, + }; + agent.emit_event(AgentEvent::ToolExecutionFinished { + result: result.clone(), + }); + results[index] = Some(CompletedToolExecution { + result, + task_succeeded: false, + should_end_turn: false, + terminated: false, + tool_name: call.name.clone(), + details: None, + }); + continue; + }; + + let ctx = self.parallel_tool_context(agent, options, &call); + if let Some(result) = self.authorize_tool_call(&call, &tool, &ctx).await? { + let execution = self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + None, + ); + results[index] = Some(execution); + continue; + } + + // Pre-execution hook check + match self.apply_pre_hooks(&mut call).await? { + None => {} + Some(reason) => { + self.emit_tool_execution_blocked(&call, &reason); + let result = ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: format!("Blocked by pre-execution hook: {reason}").into(), + is_error: true, + }; + let execution = self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + None, + ); + results[index] = Some(execution); + continue; + } + } + + if let Err(error) = self.emit_tool_runtime_started(&call) { + let result = self.blocked_tool_result(&call, error); + let execution = self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + None, + ); + results[index] = Some(execution); + continue; + } + + join_set.spawn(async move { + let output = execute_tool_future( + &call.name, + descriptor.execution_timeout, + tool.execute_output(ctx, call.input.clone()), + ) + .await; + (index, call, descriptor, output) + }); + } + + while !join_set.is_empty() { + if let Err(error) = options.check_limits() { + join_set.abort_all(); + return Err(error); + } + match tokio::time::timeout(PARALLEL_JOIN_POLL_INTERVAL, join_set.join_next()).await { + Ok(Some(Ok((index, call, descriptor, output)))) => { + let (result, details, terminate) = self.tool_output_block(&call, output).await; + // RUNTIME defense: a parallel-lane execution can never end + // the run — a `terminate: true` surfacing here is a tool + // misuse (or a static-coercion gap), never honored as + // termination, and never a silent race with the rest of + // the batch. + let (result, details) = if terminate { + eprintln!( + "warning: tool '{}' requested termination from a parallel \ + execution; rejecting as a misuse error, run continues", + call.name + ); + (parallel_termination_rejected(&call), None) + } else { + (result, details) + }; + results[index] = Some(self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + details, + )); + } + Ok(Some(Err(error))) => { + join_set.abort_all(); + return Err(RuntimeError::Store(format!( + "parallel tool task failed: {error}" + ))); + } + Ok(None) => break, + Err(_) => continue, + } + } + + if let Err(error) = options.check_limits() { + join_set.abort_all(); + return Err(error); + } + + let mut ordered = Vec::with_capacity(len); + for result in results { + ordered.push(result.ok_or_else(|| { + RuntimeError::Store("parallel tool batch lost a result".to_string()) + })?); + } + + Ok(ordered) + } + + async fn execute_registered_tool( + &mut self, + agent: &mut Agent, + options: &RunOptions, + mut call: ToolCall, + ) -> CompletedToolExecution { + let Some((tool, descriptor)) = self.registered_tool(&call.name) else { + let result = ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: "Tool not found".into(), + is_error: true, + }; + agent.emit_event(AgentEvent::ToolExecutionFinished { + result: result.clone(), + }); + return CompletedToolExecution { + result, + task_succeeded: false, + should_end_turn: false, + terminated: false, + tool_name: call.name.clone(), + details: None, + }; + }; + + let authorization_ctx = self.parallel_tool_context(agent, options, &call); + match self + .authorize_tool_call(&call, &tool, &authorization_ctx) + .await + { + Ok(Some(result)) => { + return self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + None, + ); + } + Ok(None) => {} + Err(error) => { + let result = self.blocked_tool_result(&call, error); + return self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + None, + ); + } + } + + // Pre-execution hook check + match self.apply_pre_hooks(&mut call).await { + Ok(None) => {} + Ok(Some(reason)) => { + self.emit_tool_execution_blocked(&call, &reason); + let result = ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: format!("Blocked by pre-execution hook: {reason}").into(), + is_error: true, + }; + return self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + None, + ); + } + Err(error) => { + let result = self.blocked_tool_result(&call, error); + return self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + None, + ); + } + } + + if let Err(error) = self.emit_tool_runtime_started(&call) { + let result = self.blocked_tool_result(&call, error); + return self.completed_execution( + agent, + &call, + &descriptor, + result, + RoundEffect::default(), + None, + ); + } + + let working_directory = authorization_ctx.working_directory.clone(); + let runtime = authorization_ctx.runtime.clone(); + let event_tx = agent.event_sender(); + let (result, details, terminate) = self + .tool_output_block( + &call, + execute_tool_future( + &call.name, + descriptor.execution_timeout, + tool.execute_mut_output( + ToolContext { + agent_id: self.agent_id.clone(), + tool_call_id: call.id.clone(), + tool_name: call.name.clone(), + working_directory, + runtime, + agent, + event_tx, + run_options: options.clone(), + }, + call.input.clone(), + ), + ) + .await, + ) + .await; + let effect = RoundEffect { + should_end_turn: agent.take_idle_requested() || terminate, + terminated: terminate, + }; + self.completed_execution(agent, &call, &descriptor, result, effect, details) + } +} + +impl ToolCallSchedule { + fn new(runtime: &ToolRuntime, agent: &Agent, calls: Vec) -> Self { + let mut batches = Vec::new(); + let mut pending_parallel = Vec::new(); + + for call in calls { + match runtime.call_execution_category_for_agent(&call, Some(agent)) { + ToolExecutionCategory::ReadOnlyParallel => pending_parallel.push(call), + ToolExecutionCategory::ExclusiveLocalMutation + | ToolExecutionCategory::ExclusivePersistentMutation + | ToolExecutionCategory::BackgroundJob + | ToolExecutionCategory::Delegation => { + if !pending_parallel.is_empty() { + batches.push(ToolCallBatch::Parallel(std::mem::take( + &mut pending_parallel, + ))); + } + batches.push(ToolCallBatch::Exclusive(call)); + } + } + } + + if !pending_parallel.is_empty() { + batches.push(ToolCallBatch::Parallel(pending_parallel)); + } + + Self { batches } + } +} + +impl ToolCallBatch { + fn execution_count(&self) -> usize { + match self { + ToolCallBatch::Exclusive(_) => 1, + ToolCallBatch::Parallel(calls) => calls.len(), + } + } + + /// Unwraps this batch into its constituent calls, in original call order. + /// Used to build not-executed results for batches skipped by termination. + fn into_calls(self) -> Vec { + match self { + ToolCallBatch::Exclusive(call) => vec![call], + ToolCallBatch::Parallel(calls) => calls, + } + } +} + +/// Builds the is_error result for a call that was never executed because an +/// earlier call in the same round terminated the run. +fn not_executed_result(call: &ToolCall, terminated_by: &str) -> ContentBlock { + ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: format!("not executed: run terminated by '{terminated_by}'").into(), + is_error: true, + } +} + +/// Builds the is_error result for a parallel-lane call that requested +/// termination — RUNTIME defense: never honored, always surfaced as misuse. +fn parallel_termination_rejected(call: &ToolCall) -> ContentBlock { + ContentBlock::ToolResult { + tool_use_id: call.id.clone(), + content: format!( + "not honored: tool '{}' requested termination from a parallel execution; \ + termination is only honored from an exclusive execution", + call.name + ) + .into(), + is_error: true, + } +} + +async fn execute_tool_future( + tool_name: &str, + execution_timeout: Option, + future: F, +) -> Result +where + F: Future>, +{ + match execution_timeout { + Some(timeout) => match tokio::time::timeout(timeout, future).await { + Ok(result) => result, + Err(_) => Err(format!( + "Tool '{tool_name}' timed out after {}", + format_duration(timeout) + )), + }, + None => future.await, + } +} + +fn format_duration(duration: Duration) -> String { + if duration.as_secs() > 0 && duration.subsec_nanos() == 0 { + format!("{}s", duration.as_secs()) + } else if duration.as_millis() > 0 { + format!("{}ms", duration.as_millis()) + } else if duration.as_micros() > 0 { + format!("{}us", duration.as_micros()) + } else { + format!("{}ns", duration.as_nanos()) + } +} diff --git a/vendor/mentra/src/tool/paging.rs b/vendor/mentra/src/tool/paging.rs new file mode 100644 index 0000000..1aa5050 --- /dev/null +++ b/vendor/mentra/src/tool/paging.rs @@ -0,0 +1,355 @@ +//! Windowing for oversized tool results. +//! +//! An agent loop cannot control how much a tool returns, so a single +//! oversized result can overflow the model's context before the run can +//! react. This module computes the *model's view* of such a result: the first +//! window plus a trailer naming the follow-up call that returns the next one. +//! The full text is retained separately (see [`PagedToolResults`]) so nothing +//! is lost — paging never discards, it defers. +//! +//! Line numbers in every trailer are **absolute over the full result**, so a +//! line the model quotes from window three means the same line it would have +//! meant in an unpaged result. + +use std::{ + collections::HashMap, + sync::{Arc, Mutex}, +}; + +use crate::agent::ToolResultPagingConfig; + +/// Name of the built-in tool that returns further windows. Referenced by the +/// paging trailer, so the two must always agree. +pub(crate) const READ_TOOL_RESULT_TOOL: &str = "read_tool_result"; + +/// Computes windows over a tool result according to an agent's paging +/// configuration. Holds no state: every window is derived from the full text +/// it is handed. +#[derive(Debug, Clone, Copy)] +pub(crate) struct ToolResultPager { + threshold_bytes: usize, + page_bytes: usize, +} + +/// One window of a full result, before the trailer is appended. +struct Window<'a> { + text: &'a str, + /// 1-based absolute line this window starts at; equals `last_line + 1` + /// when the window is empty because `start_line` was past the end. + first_line: usize, + /// 1-based absolute line this window ends at, or `first_line - 1` when + /// the window is empty. + last_line: usize, + total_lines: usize, + /// Set when a single line exceeded `page_bytes` and had to be cut mid-line: + /// `(line number, bytes shown, bytes in the whole line)`. + hard_cut: Option<(usize, usize, usize)>, +} + +impl ToolResultPager { + pub(crate) fn new(config: ToolResultPagingConfig) -> Self { + Self { + threshold_bytes: config.threshold_bytes, + // A zero-byte page would emit nothing but markers forever; one + // byte still guarantees forward progress line by line. + page_bytes: config.page_bytes.max(1), + } + } + + /// Returns the first window when `text` is oversized, or `None` when it + /// is at or below the threshold and must be inserted unchanged. + pub(crate) fn first_page(&self, tool_use_id: &str, text: &str) -> Option { + if text.len() <= self.threshold_bytes { + return None; + } + Some(self.window(tool_use_id, text, 1)) + } + + /// Returns the window starting at the 1-based absolute `start_line`, + /// terminated by either a paging trailer or the end-of-result marker. A + /// `start_line` past the end yields an empty window and the end marker. + pub(crate) fn window(&self, tool_use_id: &str, text: &str, start_line: usize) -> String { + let window = self.cut(text, start_line.max(1)); + let mut rendered = String::with_capacity(window.text.len() + TRAILER_HEADROOM_BYTES); + rendered.push_str(window.text); + if !rendered.is_empty() && !rendered.ends_with('\n') { + rendered.push('\n'); + } + + if let Some((line, shown, total)) = window.hard_cut { + rendered.push_str(&format!( + "…[line {} hard-cut at {} of {} bytes; the remainder of this line is skipped]\n", + thousands(line), + thousands(shown), + thousands(total), + )); + } + + if window.last_line >= window.total_lines { + rendered.push_str(END_OF_RESULT_MARKER); + return rendered; + } + + rendered.push_str(&format!( + "…[paged: lines {}–{} of {} ({} KB of {} KB). \ + Call {READ_TOOL_RESULT_TOOL}(tool_use_id=\"{tool_use_id}\", start_line={}) \ + for the next window.]", + thousands(window.first_line), + thousands(window.last_line), + thousands(window.total_lines), + kilobytes(window.text.len()), + kilobytes(text.len()), + thousands(window.last_line + 1), + )); + rendered + } + + /// Selects the slice of `text` that starts at `start_line` and fits in + /// `page_bytes`, always ending on a line boundary unless a single line is + /// itself too long to fit. + fn cut<'a>(&self, text: &'a str, start_line: usize) -> Window<'a> { + let total_lines = text.split_inclusive('\n').count(); + if start_line > total_lines { + return Window { + text: "", + first_line: start_line, + last_line: total_lines, + total_lines, + hard_cut: None, + }; + } + + let start = text + .split_inclusive('\n') + .take(start_line - 1) + .map(str::len) + .sum::(); + let mut shown = 0_usize; + let mut lines = 0_usize; + for line in text[start..].split_inclusive('\n') { + if shown + line.len() > self.page_bytes { + break; + } + shown += line.len(); + lines += 1; + } + + if lines == 0 { + // The line at `start_line` alone exceeds a whole page: the only + // case where a window may end mid-line. Cut on a character + // boundary so the window is always valid UTF-8, and resume at the + // next line — a partial line has no addressable start. + let line = text[start..] + .split_inclusive('\n') + .next() + .expect("start_line is within the result, so a line follows"); + let mut end = self.page_bytes.min(line.len()); + while end > 0 && !line.is_char_boundary(end) { + end -= 1; + } + return Window { + text: &line[..end], + first_line: start_line, + last_line: start_line, + total_lines, + hard_cut: Some((start_line, end, line.len())), + }; + } + + Window { + text: &text[start..start + shown], + first_line: start_line, + last_line: start_line + lines - 1, + total_lines, + hard_cut: None, + } + } +} + +/// Full texts of this agent's paged tool results, keyed by `tool_use_id`. +/// +/// Entries are immutable once recorded and are only ever read back whole; +/// the map lives for the life of the agent and is dropped with it. Nothing +/// here is persisted: the pager serves the live run, and the transcript +/// already records what the model actually saw. +#[derive(Clone, Default)] +pub(crate) struct PagedToolResults { + entries: Arc>>>, +} + +impl PagedToolResults { + pub(crate) fn record(&self, tool_use_id: &str, full: &str) { + self.entries + .lock() + .expect("paged tool results poisoned") + .insert(tool_use_id.to_string(), Arc::from(full)); + } + + pub(crate) fn get(&self, tool_use_id: &str) -> Option> { + self.entries + .lock() + .expect("paged tool results poisoned") + .get(tool_use_id) + .cloned() + } +} + +const END_OF_RESULT_MARKER: &str = "…[end of result]"; + +/// Slack reserved for the trailer when sizing the rendered window buffer. +const TRAILER_HEADROOM_BYTES: usize = 256; + +fn kilobytes(bytes: usize) -> String { + format!("{:.1}", bytes as f64 / 1024.0) +} + +/// Formats a count with thousands separators, matching the trailer format +/// (`lines 1–812 of 5,723`). +fn thousands(value: usize) -> String { + let digits = value.to_string(); + let mut grouped = String::with_capacity(digits.len() + digits.len() / 3); + for (index, digit) in digits.chars().enumerate() { + if index > 0 && (digits.len() - index).is_multiple_of(3) { + grouped.push(','); + } + grouped.push(digit); + } + grouped +} + +#[cfg(test)] +mod tests { + use super::*; + + fn pager(threshold_bytes: usize, page_bytes: usize) -> ToolResultPager { + ToolResultPager::new(ToolResultPagingConfig { + threshold_bytes, + page_bytes, + }) + } + + /// 26 lines of exactly 10 bytes each ("line-01xx\n" … ) = 260 bytes. + fn numbered_lines(count: usize) -> String { + (1..=count) + .map(|line| format!("line-{line:02}xx\n")) + .collect() + } + + #[test] + fn results_at_or_below_the_threshold_are_never_paged() { + let text = numbered_lines(6); + assert_eq!(text.len(), 60); + + assert_eq!(pager(60, 20).first_page("call-1", &text), None); + assert!(pager(59, 20).first_page("call-1", &text).is_some()); + } + + #[test] + fn the_first_page_carries_absolute_lines_and_byte_totals() { + let text = numbered_lines(26); + let page = pager(100, 30).first_page("call-8", &text).expect("paged"); + + assert!(page.starts_with("line-01xx\nline-02xx\nline-03xx\n")); + assert!( + page.contains( + "…[paged: lines 1–3 of 26 (0.0 KB of 0.3 KB). \ + Call read_tool_result(tool_use_id=\"call-8\", start_line=4) for the next window.]" + ), + "unexpected trailer: {page}" + ); + } + + #[test] + fn windows_tile_the_result_without_gaps_or_overlap() { + let text = numbered_lines(26); + let pager = pager(100, 30); + + let second = pager.window("call-8", &text, 4); + assert!(second.starts_with("line-04xx\nline-05xx\nline-06xx\n")); + assert!(second.contains("lines 4–6 of 26")); + assert!(second.contains("start_line=7")); + } + + #[test] + fn the_final_window_carries_the_end_marker_instead_of_a_trailer() { + let text = numbered_lines(26); + let last = pager(100, 30).window("call-8", &text, 25); + + assert!(last.starts_with("line-25xx\nline-26xx\n")); + assert!(last.ends_with("…[end of result]")); + assert!(!last.contains("[paged:")); + } + + #[test] + fn a_start_line_past_the_end_returns_an_empty_window() { + let text = numbered_lines(26); + + assert_eq!( + pager(100, 30).window("call-8", &text, 27), + "…[end of result]" + ); + assert_eq!( + pager(100, 30).window("call-8", &text, 9_999), + "…[end of result]" + ); + } + + #[test] + fn a_line_longer_than_a_page_hard_cuts_on_a_character_boundary() { + // Four-byte characters, so every page_bytes that is not a multiple of + // four must round down rather than split the character. + let text = format!("{}\nnext line\n", "𝄞".repeat(10)); + let window = pager(10, 10).window("call-8", &text, 1); + + assert!(window.starts_with("𝄞𝄞")); + assert!(!window.starts_with("𝄞𝄞𝄞")); + assert!(window.contains("…[line 1 hard-cut at 8 of 41 bytes")); + assert!(window.contains("start_line=2")); + } + + #[test] + fn the_window_after_a_hard_cut_resumes_at_the_next_whole_line() { + let text = format!("{}\nnext line\n", "𝄞".repeat(10)); + let window = pager(10, 10).window("call-8", &text, 2); + + assert!(window.starts_with("next line\n")); + assert!(window.ends_with("…[end of result]")); + } + + #[test] + fn windows_never_split_a_line_that_fits_and_preserve_crlf() { + let text = "alpha\r\nbéta\r\ngamma\r\n"; + let window = pager(4, 8).window("call-8", text, 1); + + assert!(window.starts_with("alpha\r\n")); + assert!(!window.contains("béta")); + assert!(window.contains("lines 1–1 of 3")); + } + + #[test] + fn a_result_without_a_trailing_newline_still_ends_before_the_trailer() { + let window = pager(4, 8).window("call-8", "alpha\nomega", 2); + + assert!(window.starts_with("omega\n")); + assert!(window.ends_with("…[end of result]")); + } + + #[test] + fn thousands_separators_match_the_documented_trailer_format() { + assert_eq!(thousands(0), "0"); + assert_eq!(thousands(812), "812"); + assert_eq!(thousands(5_723), "5,723"); + assert_eq!(thousands(1_234_567), "1,234,567"); + } + + #[test] + fn recorded_results_are_readable_by_tool_use_id_and_isolated_per_id() { + let store = PagedToolResults::default(); + store.record("call-1", "first"); + store.record("call-2", "second"); + + assert_eq!(store.get("call-1").as_deref(), Some("first")); + assert_eq!(store.get("call-2").as_deref(), Some("second")); + assert_eq!(store.get("call-3"), None); + } +} diff --git a/vendor/mentra/src/tool/runtime.rs b/vendor/mentra/src/tool/runtime.rs new file mode 100644 index 0000000..d8e8095 --- /dev/null +++ b/vendor/mentra/src/tool/runtime.rs @@ -0,0 +1 @@ +pub(crate) use super::orchestrator::ToolRuntime; diff --git a/vendor/mentra/src/tool/truncation.rs b/vendor/mentra/src/tool/truncation.rs new file mode 100644 index 0000000..65094af --- /dev/null +++ b/vendor/mentra/src/tool/truncation.rs @@ -0,0 +1,288 @@ +use std::{ + fs::{self, OpenOptions}, + io::Write, + path::{Path, PathBuf}, + sync::atomic::{AtomicU64, Ordering}, + time::{SystemTime, UNIX_EPOCH}, +}; + +use mentra_provider::ToolResultContent; + +static NEXT_SPILL_ID: AtomicU64 = AtomicU64::new(1); + +pub(super) enum SpillBehavior { + Enabled(PathBuf), + Disabled(&'static str), +} + +pub(super) struct ToolOutputLimiter { + max_bytes: usize, + max_lines: usize, + spill: SpillBehavior, +} + +impl ToolOutputLimiter { + pub(super) fn new(max_bytes: usize, max_lines: usize, spill: SpillBehavior) -> Self { + Self { + max_bytes, + max_lines, + spill, + } + } + + pub(super) async fn apply(&self, content: ToolResultContent) -> ToolResultContent { + match content { + ToolResultContent::Text(text) => self.apply_text(text).await, + ToolResultContent::Structured(value) => self.apply_structured(value).await, + } + } + + async fn apply_text(&self, text: String) -> ToolResultContent { + let total_lines = line_count(&text); + if text.len() <= self.max_bytes && total_lines <= self.max_lines { + return ToolResultContent::Text(text); + } + + let mut shown_bytes = 0_usize; + let mut shown_lines = 0_usize; + for line in text.split_inclusive('\n') { + if shown_lines == self.max_lines + || shown_bytes.saturating_add(line.len()) > self.max_bytes + { + break; + } + shown_bytes += line.len(); + shown_lines += 1; + } + + let mut truncated = text[..shown_bytes].to_string(); + if !truncated.is_empty() && !truncated.ends_with('\n') { + truncated.push('\n'); + } + let spill = self.spill(text, "txt").await; + truncated.push_str(&format!( + "[truncated: showing {shown_lines} of {total_lines} lines; {spill}]" + )); + ToolResultContent::Text(truncated) + } + + async fn apply_structured(&self, value: serde_json::Value) -> ToolResultContent { + let serialized = serde_json::to_string(&value) + .expect("serde_json::Value always serializes to valid JSON"); + let total_lines = line_count(&serialized); + if serialized.len() <= self.max_bytes && total_lines <= self.max_lines { + return ToolResultContent::Structured(value); + } + + let serialized_len = serialized.len(); + let spill = self.spill(serialized, "json").await; + ToolResultContent::Text(format!( + "[truncated: structured tool output is {serialized_len} bytes across {total_lines} lines; {spill}]" + )) + } + + async fn spill(&self, content: String, extension: &'static str) -> String { + match &self.spill { + SpillBehavior::Enabled(directory) => { + let directory = directory.clone(); + match tokio::task::spawn_blocking(move || { + spill_file(&directory, extension, &content) + }) + .await + { + Ok(Ok(path)) => format!("full output at {}", path.display()), + Ok(Err(error)) => format!( + "full output could not be saved ({error}); increase the tool-result limits" + ), + Err(error) => format!( + "full output could not be saved (spill task failed: {error}); increase the tool-result limits" + ), + } + } + SpillBehavior::Disabled(reason) => { + format!("full output was not saved because {reason}") + } + } + } +} + +fn line_count(text: &str) -> usize { + if text.is_empty() { + 0 + } else { + text.bytes().filter(|byte| *byte == b'\n').count() + usize::from(!text.ends_with('\n')) + } +} + +fn spill_file(directory: &Path, extension: &str, content: &str) -> Result { + fs::create_dir_all(directory).map_err(|error| { + format!( + "failed to create spill directory '{}': {error}", + directory.display() + ) + })?; + + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or_default(); + + for _ in 0..16 { + let id = NEXT_SPILL_ID.fetch_add(1, Ordering::Relaxed); + let path = directory.join(format!( + "tool-output-{}-{timestamp}-{id}.{extension}", + std::process::id() + )); + let mut options = OpenOptions::new(); + options.write(true).create_new(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options.mode(0o600); + } + + match options.open(&path) { + Ok(mut file) => { + if let Err(error) = file.write_all(content.as_bytes()) { + let _ = fs::remove_file(&path); + return Err(format!("failed to write '{}': {error}", path.display())); + } + return Ok(path); + } + Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => continue, + Err(error) => return Err(format!("failed to create '{}': {error}", path.display())), + } + } + + Err("could not allocate a unique spill filename".to_string()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn no_spill(max_bytes: usize, max_lines: usize) -> ToolOutputLimiter { + ToolOutputLimiter::new( + max_bytes, + max_lines, + SpillBehavior::Disabled("spill is disabled for this test"), + ) + } + + fn text(content: ToolResultContent) -> String { + match content { + ToolResultContent::Text(text) => text, + ToolResultContent::Structured(_) => panic!("expected text"), + } + } + + #[tokio::test] + async fn under_limit_text_is_byte_identical() { + let original = "alpha\r\nbéta\n".to_string(); + assert_eq!( + no_spill(original.len(), 2) + .apply(ToolResultContent::Text(original.clone())) + .await, + ToolResultContent::Text(original) + ); + } + + #[tokio::test] + async fn truncation_preserves_complete_crlf_and_utf8_lines() { + let result = text( + no_spill(10, 10) + .apply(ToolResultContent::Text( + "alpha\r\nbéta\r\ngamma\r\n".to_string(), + )) + .await, + ); + assert!(result.starts_with("alpha\r\n")); + assert!(!result.contains("béta")); + assert!(result.contains("showing 1 of 3 lines")); + } + + #[tokio::test] + async fn oversized_first_line_is_never_partially_emitted() { + let result = text( + no_spill(4, 10) + .apply(ToolResultContent::Text("ééé\nnext".to_string())) + .await, + ); + assert!(result.starts_with("[truncated:")); + assert!(result.contains("showing 0 of 2 lines")); + assert!(!result.contains('é')); + } + + #[tokio::test] + async fn line_limit_preserves_the_requested_head() { + let result = text( + no_spill(usize::MAX, 2) + .apply(ToolResultContent::Text("one\ntwo\nthree\n".to_string())) + .await, + ); + assert!(result.starts_with("one\ntwo\n[truncated:")); + assert!(result.contains("showing 2 of 3 lines")); + } + + #[tokio::test] + async fn structured_content_spills_whole_json_and_becomes_pointer_text() { + let directory = std::env::temp_dir().join(format!( + "mentra-tool-output-limiter-{}-{}", + std::process::id(), + NEXT_SPILL_ID.fetch_add(1, Ordering::Relaxed) + )); + let limiter = ToolOutputLimiter::new(4, 10, SpillBehavior::Enabled(directory.clone())); + let value = serde_json::json!({"answer": [1, 2, 3]}); + let pointer = text( + limiter + .apply(ToolResultContent::Structured(value.clone())) + .await, + ); + assert!(pointer.contains("structured tool output")); + assert!(pointer.contains("full output at")); + + let files = fs::read_dir(&directory) + .expect("read spill directory") + .collect::, _>>() + .expect("read spill entries"); + assert_eq!(files.len(), 1); + let stored = fs::read_to_string(files[0].path()).expect("read spill file"); + assert_eq!(stored, serde_json::to_string(&value).unwrap()); + fs::remove_dir_all(directory).expect("remove spill directory"); + } + + #[tokio::test] + async fn spill_failures_keep_text_and_structured_results_actionable() { + let blocking_file = std::env::temp_dir().join(format!( + "mentra-tool-output-blocker-{}-{}", + std::process::id(), + NEXT_SPILL_ID.fetch_add(1, Ordering::Relaxed) + )); + fs::write(&blocking_file, "not a directory").expect("create blocking file"); + let limiter = ToolOutputLimiter::new( + 4, + 10, + SpillBehavior::Enabled(blocking_file.join("tool-output")), + ); + + let text_result = text( + limiter + .apply(ToolResultContent::Text("oversized text".to_string())) + .await, + ); + assert!(text_result.contains("full output could not be saved")); + assert!(text_result.contains("increase the tool-result limits")); + + let structured_result = text( + limiter + .apply(ToolResultContent::Structured( + serde_json::json!({"oversized": true}), + )) + .await, + ); + assert!(structured_result.contains("full output could not be saved")); + assert!(structured_result.contains("increase the tool-result limits")); + assert!(blocking_file.is_file()); + fs::remove_file(blocking_file).expect("remove blocking file"); + } +} diff --git a/vendor/mentra/src/transcript.rs b/vendor/mentra/src/transcript.rs new file mode 100644 index 0000000..9330972 --- /dev/null +++ b/vendor/mentra/src/transcript.rs @@ -0,0 +1,824 @@ +use std::{ + collections::BTreeMap, + sync::atomic::{AtomicU64, Ordering}, + time::{SystemTime, UNIX_EPOCH}, +}; + +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use thiserror::Error; + +use crate::{ContentBlock, Message, Role}; + +static NEXT_ENTRY_SUFFIX: AtomicU64 = AtomicU64::new(0); + +/// Identifier for one transcript entry. +/// +/// Entries form a tree through [`TranscriptItem::parent_id`]; this is how a +/// conversation can return to an earlier point and continue along a different +/// path without copying history. +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] +pub struct EntryId(String); + +impl EntryId { + pub fn new() -> Self { + let stamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + let suffix = NEXT_ENTRY_SUFFIX.fetch_add(1, Ordering::Relaxed); + Self(format!("entry-{stamp:x}-{suffix:x}")) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl Default for EntryId { + fn default() -> Self { + Self::new() + } +} + +impl std::fmt::Display for EntryId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(&self.0) + } +} + +/// Why a branch operation could not be performed. +#[derive(Debug, Clone, PartialEq, Eq, Error)] +pub enum BranchError { + #[error("no entry '{0}' anywhere in the transcript")] + UnknownEntry(EntryId), + /// An archived entry whose parent chain does not reach a root. + /// + /// Impossible in a well-formed tree: every entry either is a root or names + /// a parent that exists. Reported rather than papered over, because + /// installing a partial path would silently hand the model a conversation + /// missing its beginning. + #[error("entry '{entry}' has a broken parent chain: '{missing}' is not in the transcript")] + BrokenChain { entry: EntryId, missing: EntryId }, +} + +/// An agent's conversation, as a tree of entries with one active path. +/// +/// [`items`](Self::items) is that active path, root to leaf — the messages +/// the model actually sees, and the only view most code needs. Entries left +/// behind by [`branch_from`](Self::branch_from) move to +/// [`archived`](Self::archived) rather than being deleted, so a branch is a +/// move of the leaf pointer rather than a copy of history — and moving it back +/// is how an abandoned branch is returned to. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(from = "AgentTranscriptWire")] +pub struct AgentTranscript { + items: Vec, + /// Entries off the active path. Reachable through + /// [`children`](Self::children), and returnable-to through + /// [`branch_from`](Self::branch_from), which accepts an archived entry and + /// rebuilds its path from the `parent_id` links. + #[serde(default, skip_serializing_if = "Vec::is_empty")] + archive: Vec, +} + +/// Deserialization shape, so transcripts written before entries had ids load +/// unchanged and get their parent links filled in on the way through. +#[derive(Deserialize)] +struct AgentTranscriptWire { + #[serde(default)] + items: Vec, + #[serde(default)] + archive: Vec, +} + +impl From for AgentTranscript { + fn from(wire: AgentTranscriptWire) -> Self { + let mut transcript = Self { + items: wire.items, + archive: wire.archive, + }; + transcript.link_active_path(); + transcript + } +} + +impl AgentTranscript { + pub fn new(items: Vec) -> Self { + let mut transcript = Self { + items, + archive: Vec::new(), + }; + transcript.link_active_path(); + transcript + } + + pub fn from_messages(messages: Vec) -> Self { + Self::new( + messages + .into_iter() + .map(transcript_item_from_message) + .collect(), + ) + } + + /// Fills in parent links the active path implies. + /// + /// The active path is a root-to-leaf chain by construction, so an entry's + /// parent is the entry before it. Only missing links are written, which + /// leaves a tree loaded from disk alone and repairs a transcript written + /// before entries had ids. + fn link_active_path(&mut self) { + for index in 1..self.items.len() { + if self.items[index].parent_id.is_none() { + self.items[index].parent_id = Some(self.items[index - 1].id.clone()); + } + } + } + + /// The entry the next append will hang from. + pub fn leaf(&self) -> Option<&EntryId> { + self.items.last().map(|item| &item.id) + } + + /// Entries that are not on the active path. + pub fn archived(&self) -> &[TranscriptItem] { + &self.archive + } + + /// Looks up an entry anywhere in the tree. + pub fn entry(&self, id: &EntryId) -> Option<&TranscriptItem> { + self.items + .iter() + .chain(self.archive.iter()) + .find(|item| &item.id == id) + } + + /// The entries recorded as continuing from `id`, in creation order. + /// + /// More than one means the conversation branched there: each is the start + /// of a different path explored from the same point. + pub fn children(&self, id: &EntryId) -> Vec<&TranscriptItem> { + self.items + .iter() + .chain(self.archive.iter()) + .filter(|item| item.parent_id.as_ref() == Some(id)) + .collect() + } + + /// Moves the leaf to `id`, so subsequent appends continue from there. + /// + /// `id` may be anywhere in the tree: on the active path, which shortens it, + /// or on a branch abandoned earlier, which returns to it. Either way no + /// entry is deleted — whatever leaves the active path moves to + /// [`archived`](Self::archived) and stays reachable through + /// [`children`](Self::children). Returns how many entries left the path. + /// + /// Returning to an abandoned branch is what makes this a tree rather than + /// an undo stack: "try something else" and "actually, go back" are the same + /// operation in opposite directions. + pub fn branch_from(&mut self, id: &EntryId) -> Result { + if let Some(position) = self.items.iter().position(|item| &item.id == id) { + let abandoned = self.items.split_off(position + 1); + let count = abandoned.len(); + self.archive.extend(abandoned); + return Ok(count); + } + + if !self.archive.iter().any(|item| &item.id == id) { + return Err(BranchError::UnknownEntry(id.clone())); + } + + // The target is on an abandoned branch. Its path is reconstructible + // because every entry names its parent, so walk to the root and make + // that chain the active path. + let path = self.path_to(id)?; + let restored: Vec = path + .iter() + .map(|id| { + self.take_anywhere(id) + .expect("path_to only names entries that exist") + }) + .collect(); + + let count = self.items.len(); + let previous = std::mem::replace(&mut self.items, restored); + self.archive.extend(previous); + + Ok(count) + } + + /// The ids from the root down to `id`, inclusive. + fn path_to(&self, id: &EntryId) -> Result, BranchError> { + let mut path = Vec::new(); + let mut cursor = Some(id.clone()); + + while let Some(current) = cursor { + let Some(item) = self.entry(¤t) else { + return Err(BranchError::BrokenChain { + entry: id.clone(), + missing: current, + }); + }; + cursor = item.parent_id.clone(); + path.push(current); + + // A parent link that cycles would loop forever. It cannot happen + // through `push`, which only ever points at an existing leaf, but + // a transcript loaded from disk is data rather than a promise. + if path.len() > self.items.len() + self.archive.len() { + return Err(BranchError::BrokenChain { + entry: id.clone(), + missing: id.clone(), + }); + } + } + + path.reverse(); + Ok(path) + } + + /// Removes an entry from whichever vector holds it. + fn take_anywhere(&mut self, id: &EntryId) -> Option { + if let Some(index) = self.items.iter().position(|item| &item.id == id) { + return Some(self.items.remove(index)); + } + let index = self.archive.iter().position(|item| &item.id == id)?; + Some(self.archive.remove(index)) + } + + pub fn items(&self) -> &[TranscriptItem] { + &self.items + } + + pub fn len(&self) -> usize { + self.items.len() + } + + pub fn is_empty(&self) -> bool { + self.items.is_empty() + } + + /// Appends an entry as a child of the current leaf. + pub fn push(&mut self, mut item: TranscriptItem) { + item.parent_id = self.leaf().cloned(); + self.items.push(item); + } + + pub fn to_messages(&self) -> Vec { + self.items + .iter() + .filter_map(TranscriptItem::project_message) + .collect() + } + + pub fn projected_messages_from(&self, start: usize) -> Vec { + self.items + .iter() + .skip(start) + .filter_map(TranscriptItem::project_message) + .collect() + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TranscriptItem { + /// Identity of this entry within the transcript tree. + #[serde(default)] + pub id: EntryId, + /// The entry this one continues from. `None` marks a root. + /// + /// Set by [`AgentTranscript::push`] rather than by the constructors: an + /// entry's parent is a property of where it is appended, not of what it + /// contains. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub parent_id: Option, + pub kind: TranscriptKind, + pub message: Option, + /// Opaque per-call host metadata attached via [`TranscriptItem::with_details`] + /// (populated from [`crate::tool::ToolOutput::details`]), keyed by + /// `tool_use_id` because one tool-result message can carry several + /// results. mentra never interprets these values; they survive + /// transcript persistence and replay but are never projected into a + /// provider request — [`TranscriptItem::project_message`] only ever + /// returns `message`. `serde(default)` keeps transcripts persisted + /// before this field existed deserializing unchanged. + #[serde(default, skip_serializing_if = "Option::is_none")] + details: Option>, +} + +impl TranscriptItem { + pub fn user_turn(message: Message) -> Self { + Self { + id: EntryId::new(), + parent_id: None, + kind: TranscriptKind::UserTurn, + message: Some(message), + details: None, + } + } + + pub fn assistant_turn(message: Message) -> Self { + Self { + id: EntryId::new(), + parent_id: None, + kind: TranscriptKind::AssistantTurn, + message: Some(message), + details: None, + } + } + + pub fn tool_exchange(message: Message, tool_use_id: Option, is_error: bool) -> Self { + Self { + id: EntryId::new(), + parent_id: None, + kind: TranscriptKind::ToolExchange { + tool_use_id, + is_error, + }, + message: Some(message), + details: None, + } + } + + pub fn canonical_context(message: Message) -> Self { + Self { + id: EntryId::new(), + parent_id: None, + kind: TranscriptKind::CanonicalContext, + message: Some(message), + details: None, + } + } + + pub fn delegation_request( + message: Message, + delegation: DelegationArtifact, + edge: Option, + ) -> Self { + Self { + id: EntryId::new(), + parent_id: None, + kind: TranscriptKind::DelegationRequest { delegation, edge }, + message: Some(message), + details: None, + } + } + + pub fn delegation_result( + message: Message, + delegation: DelegationArtifact, + edge: Option, + ) -> Self { + Self { + id: EntryId::new(), + parent_id: None, + kind: TranscriptKind::DelegationResult { delegation, edge }, + message: Some(message), + details: None, + } + } + + pub fn compaction_summary(summary: CompactionSummary) -> Self { + Self { + message: Some(Message::user(ContentBlock::text( + summary.render_for_handoff(), + ))), + id: EntryId::new(), + parent_id: None, + kind: TranscriptKind::CompactionSummary { summary }, + details: None, + } + } + + /// Attaches opaque per-call host metadata to this item, keyed by + /// `tool_use_id`. A no-op for an empty map, so attaching a possibly-empty + /// collected map never turns a details-free item into one carrying + /// `Some(empty map)`. + pub fn with_details(mut self, details: BTreeMap) -> Self { + if !details.is_empty() { + self.details = Some(details); + } + self + } + + /// This item's opaque per-call host metadata, if any. mentra never + /// interprets these values — a host recovers its own metadata after a + /// round through this accessor alone, without mentra knowing any host + /// type. + pub fn details(&self) -> Option<&BTreeMap> { + self.details.as_ref() + } + + /// Looks up this item's opaque metadata for one `tool_use_id`. + pub fn detail(&self, tool_use_id: &str) -> Option<&Value> { + self.details.as_ref()?.get(tool_use_id) + } + + pub fn project_message(&self) -> Option { + self.message.clone() + } + + pub fn is_real_user_turn(&self) -> bool { + matches!(self.kind, TranscriptKind::UserTurn) + } + + pub fn is_delegation_result(&self) -> bool { + matches!(self.kind, TranscriptKind::DelegationResult { .. }) + } + + pub fn text(&self) -> String { + self.message.as_ref().map(Message::text).unwrap_or_default() + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum TranscriptKind { + UserTurn, + AssistantTurn, + ToolExchange { + #[serde(default, skip_serializing_if = "Option::is_none")] + tool_use_id: Option, + is_error: bool, + }, + CanonicalContext, + MemoryRecall, + DelegationRequest { + delegation: DelegationArtifact, + #[serde(default, skip_serializing_if = "Option::is_none")] + edge: Option, + }, + DelegationResult { + delegation: DelegationArtifact, + #[serde(default, skip_serializing_if = "Option::is_none")] + edge: Option, + }, + CompactionSummary { + summary: CompactionSummary, + }, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum DelegationKind { + Subagent, + Teammate, + Parent, + Child, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum DelegationStatus { + Requested, + Finished, + Failed, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct DelegationEdge { + pub kind: DelegationKind, + pub local_agent_id: String, + pub remote_agent_id: String, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct DelegationArtifact { + pub kind: DelegationKind, + pub agent_id: String, + pub agent_name: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub role: Option, + pub status: DelegationStatus, + pub task_summary: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub result_summary: Option, + #[serde(default)] + pub artifacts: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct CompactionSummary { + pub goal: String, + pub progress: String, + #[serde(default)] + pub decisions: Vec, + #[serde(default)] + pub constraints: Vec, + #[serde(default)] + pub delegated_work: Vec, + #[serde(default)] + pub artifacts: Vec, + #[serde(default)] + pub open_questions: Vec, + #[serde(default)] + pub next_steps: Vec, + /// Files the agent has read or modified, accumulated across every + /// compaction in this transcript's history. + /// + /// Carried structurally rather than left to the model's prose: a summary + /// is itself summarized by the next compaction, so a file list that lived + /// only in `progress` decayed out of context after two or three rounds, + /// silently. The agent would simply stop knowing it edited something an + /// hour ago. + #[serde(default)] + pub files_touched: Vec, +} + +impl CompactionSummary { + pub fn render_for_handoff(&self) -> String { + let mut lines = vec![ + "[Compaction summary]".to_string(), + format!("Goal: {}", fallback_text(&self.goal)), + format!("Progress: {}", fallback_text(&self.progress)), + ]; + append_list(&mut lines, "Decisions", &self.decisions); + append_list(&mut lines, "Constraints", &self.constraints); + append_list(&mut lines, "Delegated work", &self.delegated_work); + append_list(&mut lines, "Artifacts", &self.artifacts); + append_list(&mut lines, "Open questions", &self.open_questions); + append_list(&mut lines, "Next steps", &self.next_steps); + append_list(&mut lines, "Files touched", &self.files_touched); + lines.join("\n") + } + + pub fn from_fallback_text(text: String) -> Self { + Self { + progress: text, + next_steps: vec![ + "Review the preserved transcript tail and continue from there.".to_string(), + ], + ..Self::default() + } + } +} + +pub(crate) fn transcript_item_from_message(message: Message) -> TranscriptItem { + match message.role { + Role::Assistant => TranscriptItem::assistant_turn(message), + Role::User => { + if let Some((tool_use_id, is_error)) = + message.content.first().and_then(|block| match block { + ContentBlock::ToolResult { + tool_use_id, + is_error, + .. + } => Some((tool_use_id.clone(), *is_error)), + _ => None, + }) + { + TranscriptItem::tool_exchange(message, Some(tool_use_id), is_error) + } else { + TranscriptItem::user_turn(message) + } + } + Role::Unknown(_) => TranscriptItem::user_turn(message), + } +} + +fn append_list(lines: &mut Vec, label: &str, items: &[String]) { + if items.is_empty() { + return; + } + lines.push(format!("{label}:")); + for item in items { + lines.push(format!("- {item}")); + } +} + +fn fallback_text(text: &str) -> &str { + if text.trim().is_empty() { + "(none)" + } else { + text + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + // Old-format compatibility (M3 test 6): a transcript persisted before + // `details` existed is exactly the JSON a details-free item serializes + // to today (the field is `skip_serializing_if` on `None`), so proving + // that JSON deserializes back to `details: None` proves genuinely old + // persisted transcripts still load. + #[test] + fn item_without_details_serializes_and_deserializes_as_old_format() { + let item = TranscriptItem::user_turn(Message::user(ContentBlock::text("hello"))); + let json = serde_json::to_string(&item).expect("serialize"); + assert!( + !json.contains("details"), + "a details-free item must serialize identically to pre-M3 transcripts, got: {json}" + ); + + let reloaded: TranscriptItem = serde_json::from_str(&json).expect("deserialize"); + assert_eq!(reloaded.details(), None); + assert_eq!(reloaded, item); + } + + #[test] + fn details_round_trip_through_json_keyed_by_tool_use_id() { + let mut details = BTreeMap::new(); + details.insert("call-1".to_string(), json!({ "secret": "shh" })); + let item = TranscriptItem::tool_exchange( + Message::user(ContentBlock::text("result")), + Some("call-1".to_string()), + false, + ) + .with_details(details.clone()); + + let json = serde_json::to_string(&item).expect("serialize"); + let reloaded: TranscriptItem = serde_json::from_str(&json).expect("deserialize"); + assert_eq!(reloaded.details(), Some(&details)); + assert_eq!(reloaded.detail("call-1"), Some(&json!({ "secret": "shh" }))); + assert_eq!(reloaded.detail("call-2"), None); + } + + #[test] + fn with_details_is_a_no_op_for_an_empty_map() { + let item = TranscriptItem::user_turn(Message::user(ContentBlock::text("hello"))) + .with_details(BTreeMap::new()); + assert_eq!(item.details(), None); + } + + // M3 test 2 (projection-boundary half): `to_messages()`/`project_message` + // are the single place internal transcript state turns into what a + // provider request carries (`Message`/`ContentBlock`) — proving details + // never appears in that projection, independent of any live agent + // plumbing, is what makes "provider requests receive only content" true + // by construction rather than by convention. The live round-trip through + // a real model request is covered by + // `agent::tests::tool_output::structured_tool_projects_content_and_hides_details_from_provider`. + #[test] + fn to_messages_projection_never_carries_details() { + let mut details = BTreeMap::new(); + details.insert("call-1".to_string(), json!({ "secret": "shh" })); + let transcript = AgentTranscript::new(vec![ + TranscriptItem::user_turn(Message::user(ContentBlock::text("go"))), + TranscriptItem::assistant_turn(Message::assistant(ContentBlock::ToolUse { + id: "call-1".to_string(), + name: "structured_details_tool".to_string(), + input: json!({}), + })), + TranscriptItem::tool_exchange( + Message::user(ContentBlock::ToolResult { + tool_use_id: "call-1".to_string(), + content: crate::tool::ToolResultContent::Structured(json!({ "answer": 42 })), + is_error: false, + }), + Some("call-1".to_string()), + false, + ) + .with_details(details), + ]); + + let projected = serde_json::to_string(&transcript.to_messages()).expect("serialize"); + assert!(projected.contains("answer"), "content must still project"); + assert!(!projected.contains("secret")); + assert!(!projected.contains("shh")); + } + + /// Builds a transcript of `n` user turns whose text is its index, so a + /// path can be described by the numbers it contains. + fn numbered(count: usize) -> AgentTranscript { + let mut transcript = AgentTranscript::default(); + for index in 0..count { + transcript.push(TranscriptItem::user_turn(Message::user( + ContentBlock::text(index.to_string()), + ))); + } + transcript + } + + /// The active path, as the numbers its entries carry. + fn path(transcript: &AgentTranscript) -> Vec { + transcript + .items() + .iter() + .filter_map(|item| item.message.as_ref()) + .map(|message| message.text()) + .collect() + } + + #[test] + fn branching_back_shortens_the_active_path() { + let mut transcript = numbered(4); + let second = transcript.items()[1].id.clone(); + + let moved = transcript.branch_from(&second).expect("branches"); + + assert_eq!(moved, 2, "two entries left the path"); + assert_eq!(path(&transcript), vec!["0", "1"]); + assert_eq!(transcript.archived().len(), 2); + } + + #[test] + fn an_abandoned_branch_can_be_returned_to() { + let mut transcript = numbered(3); + let original_leaf = transcript.leaf().expect("a leaf").clone(); + let first = transcript.items()[0].id.clone(); + + // Leave the original line of work, then explore a different one. + transcript.branch_from(&first).expect("branches away"); + transcript.push(TranscriptItem::user_turn(Message::user( + ContentBlock::text("elsewhere"), + ))); + assert_eq!(path(&transcript), vec!["0", "elsewhere"]); + + // Going back is the half that never worked: the entry is archived, so + // the old code could not find it at all. + let moved = transcript + .branch_from(&original_leaf) + .expect("returns to the abandoned branch"); + + // One, not two: entry "0" is on both paths, so only "elsewhere" + // actually left. The count is entries that stopped being active, which + // is what a caller wants to report, not the length of the old path. + assert_eq!(moved, 1, "only the entry unique to the old path left it"); + assert_eq!( + path(&transcript), + vec!["0", "1", "2"], + "the original path comes back whole and in order" + ); + } + + #[test] + fn alternating_between_two_branches_converges() { + let mut transcript = numbered(2); + let fork = transcript.items()[0].id.clone(); + let left = transcript.leaf().expect("a leaf").clone(); + + transcript.branch_from(&fork).expect("branches away"); + transcript.push(TranscriptItem::user_turn(Message::user( + ContentBlock::text("right"), + ))); + let right = transcript.leaf().expect("a leaf").clone(); + + // Three round trips: a scheme that copied entries rather than moving + // them would grow the transcript on every switch. + let total = transcript.items().len() + transcript.archived().len(); + for _ in 0..3 { + transcript.branch_from(&left).expect("goes left"); + assert_eq!(path(&transcript), vec!["0", "1"]); + transcript.branch_from(&right).expect("goes right"); + assert_eq!(path(&transcript), vec!["0", "right"]); + } + + assert_eq!( + transcript.items().len() + transcript.archived().len(), + total, + "switching branches moves entries, never copies them" + ); + } + + #[test] + fn an_unknown_entry_is_still_refused() { + let mut transcript = numbered(2); + let stranger = EntryId::new(); + + assert_eq!( + transcript.branch_from(&stranger), + Err(BranchError::UnknownEntry(stranger)) + ); + } + + #[test] + fn a_returned_to_branch_survives_a_round_trip_through_json() { + let mut transcript = numbered(3); + let leaf = transcript.leaf().expect("a leaf").clone(); + let first = transcript.items()[0].id.clone(); + + transcript.branch_from(&first).expect("branches away"); + transcript.push(TranscriptItem::user_turn(Message::user( + ContentBlock::text("elsewhere"), + ))); + + let text = serde_json::to_string(&transcript).expect("serializes"); + let mut reloaded: AgentTranscript = serde_json::from_str(&text).expect("deserializes"); + + // The archive has to survive persistence, or a branch is returnable-to + // only until the process restarts. + reloaded + .branch_from(&leaf) + .expect("a reloaded transcript can still return to its branch"); + assert_eq!(path(&reloaded), vec!["0", "1", "2"]); + } + + #[test] + fn a_child_of_an_abandoned_entry_is_still_reachable() { + let mut transcript = numbered(3); + let first = transcript.items()[0].id.clone(); + let second = transcript.items()[1].id.clone(); + + transcript.branch_from(&first).expect("branches away"); + + // `children` is what a UI offers as "you have another line of work + // here", so what it names must be what `branch_from` accepts. + let children = transcript.children(&first); + assert!(children.iter().any(|item| item.id == second)); + assert!(transcript.branch_from(&second).is_ok()); + } +} diff --git a/vendor/mentra/tests/agent_runtime.rs b/vendor/mentra/tests/agent_runtime.rs new file mode 100644 index 0000000..0996299 --- /dev/null +++ b/vendor/mentra/tests/agent_runtime.rs @@ -0,0 +1,562 @@ +use std::{ + collections::VecDeque, + fs, + sync::{ + Arc, + atomic::{AtomicU64, Ordering}, + }, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use mentra::runtime::{ProviderRetry, RunOptions, SqliteRuntimeStore}; +use mentra::{ + AgentConfig, BuiltinProvider, ContentBlock, Message, Role, Runtime, + agent::{AgentEvent, AgentStatus, RoundContext, RoundDecision, RoundStrategy}, + error::RuntimeError, + provider::{ + ContentBlockDelta, ContentBlockStart, ModelInfo, Provider, ProviderDescriptor, + ProviderError, ProviderEvent, ProviderEventStream, ProviderId, Request, + }, +}; +use reqwest::StatusCode; +use tokio::sync::{Mutex, broadcast, mpsc}; + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +enum StreamScript { + Buffered(Vec>), + Fail(ProviderError), +} + +#[derive(Clone)] +struct ScriptedProvider { + kind: ProviderId, + models: Vec, + scripts: std::sync::Arc>>, + requests: std::sync::Arc>>>, +} + +impl ScriptedProvider { + fn new( + kind: impl Into, + models: Vec, + scripts: Vec, + ) -> Self { + Self { + kind: kind.into(), + models, + scripts: std::sync::Arc::new(Mutex::new(VecDeque::from(scripts))), + requests: std::sync::Arc::new(Mutex::new(Vec::new())), + } + } + + async fn recorded_requests(&self) -> Vec> { + self.requests.lock().await.clone() + } +} + +#[async_trait] +impl Provider for ScriptedProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.kind.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(self.models.clone()) + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.requests.lock().await.push(request.into_owned()); + match self.scripts.lock().await.pop_front() { + Some(StreamScript::Buffered(items)) => { + let (tx, rx) = mpsc::unbounded_channel(); + for item in items { + tx.send(item) + .expect("test stream receiver dropped unexpectedly"); + } + Ok(rx) + } + Some(StreamScript::Fail(error)) => Err(error), + None => panic!("no scripted stream available"), + } + } +} + +#[tokio::test] +async fn send_streamed_text_turn_emits_events_and_commits_history() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![text_stream(&model.id, "Hello")], + ); + + let runtime = test_runtime(provider); + let mut agent = runtime + .spawn_with_config( + "agent", + model, + AgentConfig { + system: Some("system prompt".to_string()), + ..AgentConfig::default() + }, + ) + .unwrap(); + let mut events = agent.subscribe_events(); + + let message = agent.send(vec![ContentBlock::text("hi")]).await.unwrap(); + + assert_eq!(message, Message::assistant(ContentBlock::text("Hello"))); + assert_eq!(agent.name(), "agent"); + assert_eq!(agent.model(), "model"); + assert_eq!(agent.history().len(), 2); + assert_eq!(agent.config().system.as_deref(), Some("system prompt")); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("Hello"))) + ); + + let events = collect_events(&mut events); + assert!(events.contains(&AgentEvent::RunStarted)); + assert!(events.contains(&AgentEvent::TextDelta { + delta: "Hello".to_string(), + full_text: "Hello".to_string(), + })); + assert!(matches!(events.last(), Some(AgentEvent::RunFinished))); + + let snapshot = agent.watch_snapshot(); + assert_eq!(snapshot.borrow().status, AgentStatus::Finished); + assert_eq!(snapshot.borrow().history_len, 2); + assert!(snapshot.borrow().current_text.is_empty()); + assert!(snapshot.borrow().pending_tool_uses.is_empty()); +} + +#[tokio::test] +async fn send_failure_rolls_history_back_and_emits_run_failed() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + text_stream(&model.id, "ok"), + erroring_stream( + vec![ProviderEvent::MessageStarted { + id: "msg-2".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }], + ProviderError::MalformedStream("boom".to_string()), + ), + ], + ); + + let runtime = test_runtime(provider); + let mut agent = runtime.spawn("agent", model).unwrap(); + agent.send(vec![ContentBlock::text("first")]).await.unwrap(); + let baseline = agent.history().to_vec(); + let mut events = agent.subscribe_events(); + + let result = agent.send(vec![ContentBlock::text("second")]).await; + assert!(result.is_err()); + assert_eq!(agent.history(), baseline.as_slice()); + + let events = collect_events(&mut events); + assert!(matches!(events.last(), Some(AgentEvent::RunFailed { .. }))); +} + +#[tokio::test] +async fn send_retries_transient_provider_error_before_streaming() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + failed_request(ProviderError::Http { + status: StatusCode::SERVICE_UNAVAILABLE, + body: "offline".to_string(), + retry_after: None, + }), + text_stream(&model.id, "recovered"), + ], + ); + let provider_handle = provider.clone(); + + let runtime = test_runtime(provider); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let message = agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect("send should retry"); + + assert_eq!(message.text(), "recovered"); + assert_eq!(provider_handle.recorded_requests().await.len(), 2); + assert_eq!( + agent.last_message(), + Some(&Message::assistant(ContentBlock::text("recovered"))) + ); +} + +/// Every delay a run announced before retrying, in the order it waited them. +/// +/// `RetryAttempt` reports the delay the runner is about to take, so this is the +/// schedule as it was actually applied rather than as it was configured. +fn announced_retry_delays(events: &[AgentEvent]) -> Vec { + events + .iter() + .filter_map(|event| match event { + AgentEvent::RetryAttempt { next_delay_ms, .. } => Some(*next_delay_ms), + _ => None, + }) + .collect() +} + +/// Fails `count` times with `error`, then answers `text`. +fn failing_then_text(count: usize, error: fn() -> ProviderError, text: &str) -> Vec { + let mut scripts: Vec = (0..count).map(|_| failed_request(error())).collect(); + scripts.push(text_stream("model", text)); + scripts +} + +fn offline() -> ProviderError { + ProviderError::Http { + status: StatusCode::SERVICE_UNAVAILABLE, + body: "offline".to_string(), + retry_after: None, + } +} + +#[tokio::test(start_paused = true)] +async fn a_default_run_waits_the_schedule_mentra_has_always_waited() { + // Nobody's timing changes until they ask: 500 ms, doubling. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + failing_then_text(3, offline, "recovered"), + ); + + let runtime = test_runtime(provider); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let mut events = agent.subscribe_events(); + + let message = agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect("the run retries and then succeeds"); + + assert_eq!(message.text(), "recovered"); + assert_eq!( + announced_retry_delays(&collect_events(&mut events)), + vec![500, 1_000, 2_000] + ); +} + +#[tokio::test(start_paused = true)] +async fn a_host_schedule_replaces_the_default_delays() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + failing_then_text(3, offline, "recovered"), + ); + + let runtime = test_runtime(provider); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let mut events = agent.subscribe_events(); + + let message = agent + .run( + vec![ContentBlock::text("hello")], + RunOptions::default().with_provider_retry(ProviderRetry { + base_delay: Duration::from_secs(2), + max_delay: Duration::from_secs(6), + ..ProviderRetry::default() + }), + ) + .await + .expect("the run retries and then succeeds"); + + assert_eq!(message.text(), "recovered"); + assert_eq!( + announced_retry_delays(&collect_events(&mut events)), + vec![2_000, 4_000, 6_000], + "the host's base doubles to the host's ceiling" + ); +} + +#[tokio::test(start_paused = true)] +async fn a_rate_limit_that_names_its_window_is_waited_out() { + // The failure this exists for: a gateway answering 429 with a window far + // longer than a blip-shaped backoff would ever reach. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + failed_request(ProviderError::Http { + status: StatusCode::TOO_MANY_REQUESTS, + body: "rate limit exceeded".to_string(), + retry_after: Some(Duration::from_secs(45)), + }), + text_stream("model", "recovered"), + ], + ); + + let runtime = test_runtime(provider); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let mut events = agent.subscribe_events(); + + let message = agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect("the run waits out the window and succeeds"); + + assert_eq!(message.text(), "recovered"); + assert_eq!( + announced_retry_delays(&collect_events(&mut events)), + vec![45_000], + "the provider's window, not the schedule's 500 ms" + ); +} + +#[tokio::test(start_paused = true)] +async fn a_server_cannot_park_a_run_for_an_hour() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + failed_request(ProviderError::Http { + status: StatusCode::SERVICE_UNAVAILABLE, + body: "come back later".to_string(), + retry_after: Some(Duration::from_secs(3_600)), + }), + text_stream("model", "recovered"), + ], + ); + + let runtime = test_runtime(provider); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + let mut events = agent.subscribe_events(); + + agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect("the run succeeds after the clamped wait"); + + assert_eq!( + announced_retry_delays(&collect_events(&mut events)), + vec![60_000], + "clamped to the default one-minute ceiling" + ); +} + +/// Counters a [`RoundStrategy`] observed at the most recent round boundary, +/// captured by [`CountingStrategy`]. +#[derive(Clone, Copy, Default)] +struct RoundCounters { + rounds_completed: usize, + model_requests: usize, + transport_retries: usize, +} + +/// A [`RoundStrategy`] that records [`RoundContext`]'s counters at each boundary +/// it observes, always proceeding. +struct CountingStrategy { + last: Mutex>, +} + +impl CountingStrategy { + fn new() -> Arc { + Arc::new(Self { + last: Mutex::new(None), + }) + } + + async fn last_counters(&self) -> RoundCounters { + self.last.lock().await.expect("strategy observed a round") + } +} + +#[async_trait] +impl RoundStrategy for CountingStrategy { + async fn on_round(&self, ctx: RoundContext<'_>) -> RoundDecision { + *self.last.lock().await = Some(RoundCounters { + rounds_completed: ctx.rounds_completed(), + model_requests: ctx.model_requests(), + transport_retries: ctx.transport_retries(), + }); + RoundDecision::proceed() + } +} + +#[tokio::test] +async fn retry_and_round_counters_are_reported_distinctly() { + // A connection-open retry, then a success: `model_requests` (today's + // request-including-retries counter) must stay at 2, `rounds_completed` (the + // logical-round counter) must land at 1 — the round that needed a retry still + // counts as exactly one completed round — and the new `transport_retries` + // counter must isolate the one retry, distinct from both. + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + failed_request(ProviderError::Http { + status: StatusCode::SERVICE_UNAVAILABLE, + body: "offline".to_string(), + retry_after: None, + }), + text_stream(&model.id, "recovered"), + ], + ); + let provider_handle = provider.clone(); + let runtime = test_runtime(provider); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let strategy = CountingStrategy::new(); + let message = agent + .run( + vec![ContentBlock::text("hello")], + RunOptions::default().with_round_strategy(strategy.clone()), + ) + .await + .expect("run should retry then succeed"); + + assert_eq!(message.text(), "recovered"); + assert_eq!(provider_handle.recorded_requests().await.len(), 2); + + let counters = strategy.last_counters().await; + assert_eq!(counters.rounds_completed, 1, "one logical round completed"); + assert_eq!( + counters.model_requests, 2, + "model_requests keeps today's semantics: it counts the retry and the success" + ); + assert_eq!( + counters.transport_retries, 1, + "exactly one transient retry, isolated from rounds_completed and reported distinctly" + ); +} + +#[tokio::test] +async fn resume_replays_last_failed_turn() { + let model = model_info("model", BuiltinProvider::Anthropic); + let provider = ScriptedProvider::new( + BuiltinProvider::Anthropic, + vec![model.clone()], + vec![ + erroring_stream( + vec![ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: model.id.clone(), + role: Role::Assistant, + }], + ProviderError::MalformedStream("boom".to_string()), + ), + text_stream(&model.id, "done"), + ], + ); + + let runtime = test_runtime(provider); + let mut agent = runtime.spawn("agent", model).expect("spawn agent"); + + let error = agent + .send(vec![ContentBlock::text("retry me")]) + .await + .expect_err("first send should fail"); + assert!(matches!(error, RuntimeError::FailedToStreamResponse(_))); + assert!(agent.history().is_empty()); + + let resumed = agent + .resume() + .await + .expect("resume should replay user turn"); + + assert_eq!(resumed.text(), "done"); + assert_eq!(agent.history().len(), 2); + assert_eq!( + agent.history()[0], + Message::user(ContentBlock::text("retry me")) + ); + assert_eq!( + agent.history()[1], + Message::assistant(ContentBlock::text("done")) + ); + + let error = agent + .resume() + .await + .expect_err("successful run clears resume state"); + assert!(matches!(error, RuntimeError::NoResumableTurn)); +} + +fn collect_events(receiver: &mut broadcast::Receiver) -> Vec { + let mut events = Vec::new(); + while let Ok(event) = receiver.try_recv() { + events.push(event); + } + events +} + +fn test_runtime(provider: ScriptedProvider) -> Runtime { + Runtime::empty_builder() + .with_provider_instance(provider) + .with_store(temp_store("agent-runtime")) + .build() + .expect("build runtime") +} + +fn temp_store(label: &str) -> SqliteRuntimeStore { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = std::env::temp_dir().join(format!( + "mentra-agent-runtime-{label}-{timestamp}-{unique}.sqlite" + )); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).expect("create temp dir"); + } + SqliteRuntimeStore::new(path) +} + +fn model_info(id: &str, provider: impl Into) -> ModelInfo { + ModelInfo::new(id, provider) +} + +fn buffered_stream(events: Vec) -> StreamScript { + StreamScript::Buffered(events.into_iter().map(Ok).collect()) +} + +fn erroring_stream(mut events: Vec, error: ProviderError) -> StreamScript { + let mut items = events.drain(..).map(Ok).collect::>(); + items.push(Err(error)); + StreamScript::Buffered(items) +} + +fn failed_request(error: ProviderError) -> StreamScript { + StreamScript::Fail(error) +} + +fn text_stream(model: &str, text: &str) -> StreamScript { + buffered_stream(vec![ + ProviderEvent::MessageStarted { + id: "msg-text".to_string(), + model: model.to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text(text.to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ]) +} diff --git a/vendor/mentra/tests/branching.rs b/vendor/mentra/tests/branching.rs new file mode 100644 index 0000000..393be5c --- /dev/null +++ b/vendor/mentra/tests/branching.rs @@ -0,0 +1,175 @@ +//! Public-API tests for the transcript entry tree. +//! +//! Branching is what lets a conversation return to an earlier point and +//! continue differently — "undo that exchange and try another instruction", +//! editing a message and re-running, or exploring two approaches from a shared +//! prefix — without starting a new session and replaying a prefix by hand. + +use mentra::{ + AgentTranscript, ContentBlock, Message, + transcript::{EntryId, TranscriptItem}, +}; + +fn user(text: &str) -> TranscriptItem { + TranscriptItem::user_turn(Message::user(ContentBlock::text(text))) +} + +fn assistant(text: &str) -> TranscriptItem { + TranscriptItem::assistant_turn(Message::assistant(ContentBlock::text(text))) +} + +fn transcript_of(texts: &[&str]) -> AgentTranscript { + let mut transcript = AgentTranscript::default(); + for (index, text) in texts.iter().enumerate() { + if index % 2 == 0 { + transcript.push(user(text)); + } else { + transcript.push(assistant(text)); + } + } + transcript +} + +fn texts(transcript: &AgentTranscript) -> Vec { + transcript + .items() + .iter() + .map(TranscriptItem::text) + .collect() +} + +#[test] +fn appending_hangs_each_entry_from_the_leaf() { + let transcript = transcript_of(&["one", "two", "three"]); + let items = transcript.items(); + + assert_eq!(items[0].parent_id, None, "the first entry is a root"); + assert_eq!(items[1].parent_id.as_ref(), Some(&items[0].id)); + assert_eq!(items[2].parent_id.as_ref(), Some(&items[1].id)); + assert_eq!(transcript.leaf(), Some(&items[2].id)); +} + +#[test] +fn branching_rewinds_the_active_path_without_deleting_anything() { + let mut transcript = transcript_of(&["ask", "answer", "follow up"]); + let rewind_to = transcript.items()[0].id.clone(); + + let moved = transcript + .branch_from(&rewind_to) + .expect("entry is on path"); + + assert_eq!(moved, 2, "two entries leave the active path"); + assert_eq!(texts(&transcript), vec!["ask"]); + assert_eq!(transcript.leaf(), Some(&rewind_to)); + assert_eq!( + transcript.archived().len(), + 2, + "abandoned entries stay in the transcript" + ); +} + +#[test] +fn a_new_turn_after_branching_becomes_a_sibling() { + let mut transcript = transcript_of(&["ask", "first answer"]); + let fork_point = transcript.items()[0].id.clone(); + + transcript.branch_from(&fork_point).expect("on path"); + transcript.push(assistant("second answer")); + + let children = transcript.children(&fork_point); + assert_eq!(children.len(), 2, "the fork point now has two paths"); + + let child_texts: Vec = children.iter().map(|item| item.text()).collect(); + assert!(child_texts.contains(&"first answer".to_string())); + assert!(child_texts.contains(&"second answer".to_string())); + + // Only the new path is live. + assert_eq!(texts(&transcript), vec!["ask", "second answer"]); +} + +#[test] +fn an_abandoned_branch_can_be_returned_to() { + let mut transcript = transcript_of(&["ask", "first answer"]); + let fork_point = transcript.items()[0].id.clone(); + let first_answer = transcript.items()[1].id.clone(); + + transcript.branch_from(&fork_point).expect("on path"); + transcript.push(assistant("second answer")); + + // The abandoned entry is still addressable, which is what makes this a + // branch rather than a truncation. + let recovered = transcript.entry(&first_answer).expect("still present"); + assert_eq!(recovered.text(), "first answer"); +} + +#[test] +fn branching_to_an_unknown_entry_is_refused() { + let mut transcript = transcript_of(&["ask"]); + + let error = transcript + .branch_from(&EntryId::new()) + .expect_err("an entry that was never appended is not a branch point"); + + assert!(error.to_string().contains("no entry")); +} + +#[test] +fn branching_to_the_leaf_changes_nothing() { + let mut transcript = transcript_of(&["ask", "answer"]); + let leaf = transcript.leaf().cloned().expect("a leaf"); + + let moved = transcript.branch_from(&leaf).expect("the leaf is on path"); + + assert_eq!(moved, 0); + assert_eq!(texts(&transcript), vec!["ask", "answer"]); + assert!(transcript.archived().is_empty()); +} + +#[test] +fn the_tree_survives_a_round_trip() { + let mut transcript = transcript_of(&["ask", "first answer"]); + let fork_point = transcript.items()[0].id.clone(); + transcript.branch_from(&fork_point).expect("on path"); + transcript.push(assistant("second answer")); + + let encoded = serde_json::to_string(&transcript).expect("serializes"); + let decoded: AgentTranscript = serde_json::from_str(&encoded).expect("deserializes"); + + assert_eq!(decoded, transcript); + assert_eq!( + decoded.children(&fork_point).len(), + 2, + "both paths survive persistence" + ); +} + +#[test] +fn a_transcript_written_before_entries_had_ids_still_loads_linked() { + // Build the shape mentra persisted previously by stripping the tree + // fields back out of a current transcript, so the fixture cannot drift + // from the real serialization format. + let modern = transcript_of(&["ask", "answer"]); + let mut encoded: serde_json::Value = serde_json::to_value(&modern).expect("serializes"); + for item in encoded["items"] + .as_array_mut() + .expect("items is an array") + .iter_mut() + { + let object = item.as_object_mut().expect("each item is an object"); + object.remove("id"); + object.remove("parent_id"); + } + + let transcript: AgentTranscript = + serde_json::from_value(encoded).expect("a pre-tree transcript still deserializes"); + let items = transcript.items(); + + assert_eq!(items.len(), 2); + assert_eq!(texts(&transcript), vec!["ask", "answer"]); + assert_eq!(items[0].parent_id, None, "the first entry is a root"); + assert_eq!( + items[1].parent_id.as_ref(), + Some(&items[0].id), + "migration links the chain, so a legacy transcript can be branched" + ); +} diff --git a/vendor/mentra/tests/mcp_sse_smoke.rs b/vendor/mentra/tests/mcp_sse_smoke.rs new file mode 100644 index 0000000..09b1843 --- /dev/null +++ b/vendor/mentra/tests/mcp_sse_smoke.rs @@ -0,0 +1,54 @@ +//! Optional manual smoke test against a real MCP HTTP+SSE server. +//! +//! This is ignored by default and never runs in ordinary CI: it needs a live +//! endpoint, which no automated run should depend on. The transport itself is +//! covered by the deterministic fixture tests in `mentra::mcp::sse`. +//! +//! Point it at any server speaking the 2024-11-05 HTTP+SSE transport: +//! +//! ```text +//! MENTRA_MCP_SSE_URL=https://mcp.example.com/sse \ +//! MENTRA_MCP_SSE_TOKEN= \ +//! cargo test -p mentra --test mcp_sse_smoke -- --ignored --nocapture +//! ``` +//! +//! `MENTRA_MCP_SSE_TOKEN` is optional. The test performs only `initialize` and +//! `tools/list`; it never calls a tool, so it cannot cause a side effect on the +//! server it is pointed at. + +use mentra::{McpSseClient, McpSseServerConfig}; + +/// Environment variable naming the SSE endpoint to probe. +const URL_VAR: &str = "MENTRA_MCP_SSE_URL"; +/// Environment variable carrying an optional bearer token. +const TOKEN_VAR: &str = "MENTRA_MCP_SSE_TOKEN"; + +#[tokio::test] +#[ignore = "requires a live MCP server; set MENTRA_MCP_SSE_URL"] +async fn initializes_and_lists_tools_against_a_live_server() { + let url = std::env::var(URL_VAR) + .unwrap_or_else(|_| panic!("set {URL_VAR} to the server's SSE endpoint")); + + let mut config = McpSseServerConfig::new("smoke", &url); + if let Ok(token) = std::env::var(TOKEN_VAR) { + config = config.with_bearer_token(token); + } + + let client = McpSseClient::connect(&config) + .await + .expect("the handshake should complete"); + + let info = client + .server_info() + .expect("initialize should report server info"); + println!("connected to {} {}", info.name, info.version); + + println!("{} tools advertised:", client.tools().len()); + for tool in client.tools() { + println!(" {}", tool.name); + } + + // No tool is called: a smoke test must not cause side effects on whatever + // server it happens to be pointed at. + client.shutdown().await; +} diff --git a/vendor/mentra/tests/public_api.rs b/vendor/mentra/tests/public_api.rs new file mode 100644 index 0000000..0033f86 --- /dev/null +++ b/vendor/mentra/tests/public_api.rs @@ -0,0 +1,859 @@ +use std::{ + collections::VecDeque, + io::{Read, Write}, + net::TcpListener, + sync::{Arc, Mutex}, + thread, + time::{SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use mentra::{ + Agent, BuiltinProvider, ContentBlock, FileToolProfile, ModelInfo, ModelSelector, Runtime, + error::RuntimeError, + provider::{ + Provider, ProviderDescriptor, ProviderError, ProviderEventStream, ProviderId, Request, + Response, Role, provider_event_stream_from_response, + }, + runtime::{ + CommandOutput, CommandRequest, RuntimeExecutor, RuntimePolicy, VolatileRuntimeStore, + }, + tool::{ParallelToolContext, ToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSpec}, +}; +use serde_json::{Value, json}; + +#[derive(Debug)] +enum Turn { + Text(String), + ToolCalls(Vec), +} + +#[derive(Debug, Clone)] +struct ScriptedToolCall { + id: Option, + name: String, + input: Value, +} + +impl ScriptedToolCall { + fn new(name: impl Into, input: Value) -> Self { + Self { + id: None, + name: name.into(), + input, + } + } +} + +#[derive(Clone)] +struct ScriptedProvider { + kind: ProviderId, + models: Vec, + turns: Arc>>, + requests: Arc>>>, +} + +impl ScriptedProvider { + fn new(kind: ProviderId, models: Vec) -> Self { + Self { + kind, + models, + turns: Arc::new(Mutex::new(VecDeque::new())), + requests: Arc::new(Mutex::new(Vec::new())), + } + } + + fn push_turns(&self, turns: Vec) { + let mut queue = self.turns.lock().expect("scripted turn queue poisoned"); + queue.extend(turns); + } + + fn recorded_requests(&self) -> Vec> { + self.requests + .lock() + .expect("scripted request log poisoned") + .clone() + } +} + +#[async_trait] +impl Provider for ScriptedProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.kind.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(self.models.clone()) + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.requests + .lock() + .expect("scripted request log poisoned") + .push(request.into_owned()); + + let turn = self + .turns + .lock() + .expect("scripted turn queue poisoned") + .pop_front() + .unwrap_or_else(|| panic!("no scripted turn remaining for public API test")); + + match turn { + Turn::Text(text) => Ok(provider_event_stream_from_response(Response { + id: format!("public-response-{}", now_nanos()), + model: self.models[0].id.clone(), + role: Role::Assistant, + content: vec![ContentBlock::text(text)], + stop_reason: None, + usage: None, + })), + Turn::ToolCalls(calls) => Ok(provider_event_stream_from_response(Response { + id: format!("public-response-{}", now_nanos()), + model: self.models[0].id.clone(), + role: Role::Assistant, + content: calls + .into_iter() + .enumerate() + .map(|(index, call)| ContentBlock::ToolUse { + id: call.id.unwrap_or_else(|| format!("tool-{}", index + 1)), + name: call.name, + input: call.input, + }) + .collect(), + stop_reason: Some("tool_use".to_string()), + usage: None, + })), + } + } +} + +struct Harness { + runtime: Runtime, + provider: ScriptedProvider, + model: ModelInfo, +} + +impl Harness { + fn new(turns: Vec) -> Self { + let runtime_id = format!("public-api-{}", now_nanos()); + let model = ModelInfo::new("mock-model", BuiltinProvider::OpenAI); + let provider = ScriptedProvider::new(model.provider.clone(), vec![model.clone()]); + provider.push_turns(turns); + + let runtime = Runtime::builder() + .with_runtime_identifier(runtime_id) + .with_store(VolatileRuntimeStore::new()) + .with_provider_instance(provider.clone()) + .build() + .expect("build runtime"); + + Self { + runtime, + provider, + model, + } + } + + fn spawn(&self, name: &str) -> Agent { + self.runtime + .spawn(name, self.model.clone()) + .expect("spawn test agent") + } + + async fn recorded_requests(&self) -> Vec> { + self.provider.recorded_requests() + } +} + +struct EchoTool; + +struct AlphaTool; + +struct EndTurnTool; + +struct SubagentSummaryTool; + +#[async_trait] +impl ToolDefinition for EchoTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("echo_tool") + .description("Echo a canned result") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .build() + } +} + +#[async_trait] +impl ToolExecutor for EchoTool { + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + Ok("echoed".to_string()) + } +} + +#[async_trait] +impl ToolDefinition for AlphaTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("alpha_tool") + .description("Return a canned alpha result") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .build() + } +} + +#[async_trait] +impl ToolExecutor for AlphaTool { + async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult { + Ok("alpha".to_string()) + } +} + +#[async_trait] +impl ToolDefinition for EndTurnTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("stop_here") + .description("End the current turn without a follow-up assistant message") + .input_schema(json!({ + "type": "object", + "properties": {} + })) + .build() + } +} + +#[async_trait] +impl ToolExecutor for EndTurnTool { + async fn execute_mut(&self, mut ctx: ToolContext<'_>, _input: Value) -> ToolResult { + ctx.request_idle(); + Ok("stopping now".to_string()) + } +} + +#[async_trait] +impl ToolDefinition for SubagentSummaryTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("subagent_summary") + .description("Spawn a disposable subagent and return its summary") + .input_schema(json!({ + "type": "object", + "properties": { + "prompt": { "type": "string" } + }, + "required": ["prompt"] + })) + .build() + } +} + +#[async_trait] +impl ToolExecutor for SubagentSummaryTool { + async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult { + let prompt = input + .get("prompt") + .and_then(|value| value.as_str()) + .ok_or_else(|| "prompt is required".to_string())?; + let mut child = ctx.spawn_subagent().map_err(|error| error.to_string())?; + // The public pattern for a custom tool that spawns work: the child runs + // under the parent run's derived bounds, not a fresh unbounded set. + let message = child + .run(vec![ContentBlock::text(prompt)], ctx.child_run_options()) + .await + .map_err(|error| format!("child failed: {error}"))?; + Ok(message.text()) + } +} + +#[tokio::test] +async fn send_returns_final_message_after_tool_execution() { + let harness = Harness::new(vec![ + Turn::ToolCalls(vec![ScriptedToolCall::new("echo_tool", json!({}))]), + Turn::Text("done".to_string()), + ]); + harness.runtime.register_tool(EchoTool); + let mut agent = harness.spawn("tool-agent"); + + let message = agent + .send(vec![ContentBlock::text("run the tool")]) + .await + .unwrap(); + + assert_eq!(message.role, Role::Assistant); + assert_eq!(message.text(), "done"); + assert_eq!(harness.recorded_requests().await.len(), 2); +} + +#[tokio::test] +async fn runtime_exposes_registered_tool_descriptors() { + let runtime_id = format!("public-api-{}", now_nanos()); + let model = ModelInfo::new("mock-model", BuiltinProvider::OpenAI); + let provider = ScriptedProvider::new(model.provider.clone(), vec![model.clone()]); + + let runtime = Runtime::empty_builder() + .with_runtime_identifier(runtime_id) + .with_store(VolatileRuntimeStore::new()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + runtime.register_tool(EchoTool); + runtime.register_tool(AlphaTool); + + assert_eq!( + runtime.tools(), + vec![AlphaTool.descriptor(), EchoTool.descriptor()] + ); + assert_eq!( + runtime.tool_descriptor("echo_tool"), + Some(EchoTool.descriptor()) + ); + assert_eq!( + runtime.tool_descriptor("alpha_tool"), + Some(AlphaTool.descriptor()) + ); + assert_eq!(runtime.tool_descriptor("missing_tool"), None); +} + +#[test] +fn runtime_builder_publicly_selects_split_file_tools() { + let model = ModelInfo::new("mock-model", BuiltinProvider::OpenAI); + let provider = ScriptedProvider::new(model.provider.clone(), vec![model]); + let runtime = Runtime::builder() + .with_file_tools(FileToolProfile::Split) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + let names = runtime + .tools() + .into_iter() + .map(|tool| tool.provider.name) + .collect::>(); + + for name in ["read", "ls", "grep", "glob", "write", "edit"] { + assert!(names.contains(name), "missing split tool {name}"); + } + assert!(!names.contains("files")); +} + +#[tokio::test] +async fn parallel_tool_context_can_spawn_subagents_from_public_api() { + let harness = Harness::new(vec![ + Turn::ToolCalls(vec![ScriptedToolCall::new( + "subagent_summary", + json!({ "prompt": "summarize the delegated work" }), + )]), + Turn::Text("child summary".to_string()), + Turn::Text("parent complete".to_string()), + ]); + harness.runtime.register_tool(SubagentSummaryTool); + let mut agent = harness.spawn("parent-agent"); + + let message = agent + .send(vec![ContentBlock::text("delegate that")]) + .await + .unwrap(); + + assert_eq!(message.role, Role::Assistant); + assert_eq!(message.text(), "parent complete"); + assert_eq!(harness.recorded_requests().await.len(), 3); +} + +#[tokio::test] +async fn empty_assistant_response_preserves_committed_tool_results() { + let harness = Harness::new(vec![Turn::ToolCalls(vec![ScriptedToolCall::new( + "stop_here", + json!({}), + )])]); + harness.runtime.register_tool(EndTurnTool); + let mut agent = harness.spawn("idle-agent"); + + let error = agent + .send(vec![ContentBlock::text("stop after the tool")]) + .await + .unwrap_err(); + + assert!(matches!(error, RuntimeError::EmptyAssistantResponse)); + assert_eq!(harness.recorded_requests().await.len(), 1); + assert_eq!(agent.history().len(), 3); + match &agent.history()[2].content[0] { + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => { + assert_eq!(tool_use_id, "tool-1"); + assert_eq!(content, "stopping now"); + assert!(!is_error); + } + other => panic!("expected tool result block, found {other:?}"), + } +} + +#[tokio::test] +async fn resolve_model_returns_explicit_id_without_listing_models() { + let runtime_id = format!("public-api-{}", now_nanos()); + let provider = FailingListModelsProvider { + kind: BuiltinProvider::Anthropic.into(), + }; + + let runtime = Runtime::builder() + .with_runtime_identifier(runtime_id) + .with_store(VolatileRuntimeStore::new()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + + let model = runtime + .resolve_model( + BuiltinProvider::Anthropic, + ModelSelector::Id("claude-custom".to_string()), + ) + .await + .expect("resolve explicit model"); + + assert_eq!( + model, + ModelInfo::new("claude-custom", BuiltinProvider::Anthropic) + ); +} + +#[tokio::test] +async fn resolve_model_selects_newest_available_then_breaks_ties_by_id() { + let runtime_id = format!("public-api-{}", now_nanos()); + let provider = ModelListingProvider { + kind: BuiltinProvider::OpenAI.into(), + models: vec![ + model_with_created_at("zeta", BuiltinProvider::OpenAI, 1_700_000_100), + model_with_created_at("alpha", BuiltinProvider::OpenAI, 1_700_000_100), + model_with_created_at("older", BuiltinProvider::OpenAI, 1_700_000_000), + ], + }; + + let runtime = Runtime::builder() + .with_runtime_identifier(runtime_id) + .with_store(VolatileRuntimeStore::new()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + + let model = runtime + .resolve_model(BuiltinProvider::OpenAI, ModelSelector::NewestAvailable) + .await + .expect("resolve newest model"); + + assert_eq!( + model, + model_with_created_at("alpha", BuiltinProvider::OpenAI, 1_700_000_100) + ); +} + +#[tokio::test] +async fn resolve_model_reports_empty_provider_listing() { + let runtime_id = format!("public-api-{}", now_nanos()); + let provider = ModelListingProvider { + kind: BuiltinProvider::Gemini.into(), + models: Vec::new(), + }; + + let runtime = Runtime::builder() + .with_runtime_identifier(runtime_id) + .with_store(VolatileRuntimeStore::new()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + + let error = runtime + .resolve_model(BuiltinProvider::Gemini, ModelSelector::NewestAvailable) + .await + .expect_err("empty listing should fail"); + + assert!(matches!( + error, + RuntimeError::NoModelsAvailable(provider) if provider == BuiltinProvider::Gemini.into() + )); +} + +#[tokio::test] +async fn resolve_model_supports_openrouter_provider() { + let runtime_id = format!("public-api-{}", now_nanos()); + let provider = ModelListingProvider { + kind: BuiltinProvider::OpenRouter.into(), + models: vec![model_with_created_at( + "openai/gpt-4.1-mini", + BuiltinProvider::OpenRouter, + 1_741_049_700, + )], + }; + + let runtime = Runtime::builder() + .with_runtime_identifier(runtime_id) + .with_store(VolatileRuntimeStore::new()) + .with_provider_instance(provider) + .build() + .expect("build runtime"); + + let model = runtime + .resolve_model(BuiltinProvider::OpenRouter, ModelSelector::NewestAvailable) + .await + .expect("resolve newest model"); + + assert_eq!( + model, + model_with_created_at( + "openai/gpt-4.1-mini", + BuiltinProvider::OpenRouter, + 1_741_049_700, + ) + ); +} + +#[tokio::test] +async fn resolve_model_supports_ollama_provider_registration() { + let runtime = Runtime::empty_builder() + .with_ollama() + .build() + .expect("build runtime"); + + let model = runtime + .resolve_model( + BuiltinProvider::Ollama, + ModelSelector::Id("qwen2.5-coder".to_string()), + ) + .await + .expect("resolve explicit model"); + + assert_eq!( + model, + ModelInfo::new("qwen2.5-coder", BuiltinProvider::Ollama) + ); +} + +#[tokio::test] +async fn resolve_model_supports_lmstudio_provider_registration() { + let runtime = Runtime::empty_builder() + .with_lmstudio() + .build() + .expect("build runtime"); + + let model = runtime + .resolve_model( + BuiltinProvider::LmStudio, + ModelSelector::Id("local-model".to_string()), + ) + .await + .expect("resolve explicit model"); + + assert_eq!( + model, + ModelInfo::new("local-model", BuiltinProvider::LmStudio) + ); +} + +#[tokio::test] +async fn resolve_model_reports_missing_provider() { + let harness = Harness::new(vec![Turn::Text("unused".to_string())]); + + let error = harness + .runtime + .resolve_model( + BuiltinProvider::Gemini, + ModelSelector::Id("gemini-2.5-pro".to_string()), + ) + .await + .expect_err("missing provider should fail"); + + assert!(matches!( + error, + RuntimeError::ProviderNotFound(Some(provider)) + if provider == BuiltinProvider::Gemini.into() + )); +} + +#[tokio::test] +async fn runtime_accepts_provider_core_openai_compatible_instances() { + let provider_id = ProviderId::new("custom-openai-compatible"); + let (base_url, handle) = spawn_models_server( + r#"{"data":[{"id":"compat-model","name":"Compat Model","created":1}]}"#, + ); + + let mut definition = mentra::provider_core::responses::openai_definition(); + definition.descriptor.id = provider_id.clone(); + definition.descriptor.display_name = Some("Custom OpenAI-Compatible".to_string()); + definition.base_url = Some(base_url); + + let runtime = Runtime::empty_builder() + .with_registered_provider(mentra::provider_core::responses::ResponsesProvider::new( + definition, + mentra::provider_core::StaticCredentialSource::new("test-key"), + )) + .build() + .expect("build runtime"); + + let model = runtime + .resolve_model(provider_id.clone(), ModelSelector::NewestAvailable) + .await + .expect("resolve model from provider-core instance"); + + assert_eq!(model.provider, provider_id); + assert_eq!(model.id, "compat-model"); + + let captured = handle.join().expect("capture request"); + let captured_lower = captured.to_ascii_lowercase(); + assert!(captured.starts_with("GET /v1/models HTTP/1.1\r\n")); + assert!(captured_lower.contains("authorization: bearer test-key\r\n")); +} + +#[derive(Clone)] +struct FailingListModelsProvider { + kind: ProviderId, +} + +#[async_trait] +impl Provider for FailingListModelsProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.kind.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Err(ProviderError::InvalidResponse( + "list_models should not be called".to_string(), + )) + } + + async fn stream(&self, _request: Request<'_>) -> Result { + let (_tx, rx) = tokio::sync::mpsc::unbounded_channel(); + Ok(rx) + } +} + +#[derive(Clone)] +struct ModelListingProvider { + kind: ProviderId, + models: Vec, +} + +#[async_trait] +impl Provider for ModelListingProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor::new(self.kind.clone()) + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(self.models.clone()) + } + + async fn stream(&self, _request: Request<'_>) -> Result { + let (_tx, rx) = tokio::sync::mpsc::unbounded_channel(); + Ok(rx) + } +} + +/// The builder was `pub` inside a private `mod builder`, re-exported nowhere: +/// `Runtime::builder()` worked on inference, but a downstream helper taking or +/// returning a half-built runtime could not write its signature at all. This +/// test is that helper, written from outside the crate — it pins the re-export +/// at both public paths, because compiling is the whole claim. +#[test] +fn a_runtime_builder_is_a_type_downstream_code_can_name() { + fn with_volatile_store(builder: mentra::RuntimeBuilder) -> mentra::runtime::RuntimeBuilder { + builder.with_store(VolatileRuntimeStore::new()) + } + + let _ = with_volatile_store(Runtime::builder()); +} + +/// The two paths a downstream crate re-exports these from, written from +/// outside mentra so a rename is a failing test rather than a broken host. +/// +/// basis re-exports `ProviderRetry` and `ResponsesTransport` from its own +/// `basis::runtime` so a host of *its* never names mentra, which makes both +/// paths part of this crate's contract rather than an implementation detail +/// that happens to be reachable. `ResponsesTransport` is visible at the crate +/// root too; `mentra::provider::` is the one to depend on, since it sits with +/// `ResponsesRequestOptions` and the rest of the wire vocabulary. +#[test] +fn the_retry_schedule_and_transport_are_types_downstream_code_can_name() { + fn patient(retry: mentra::runtime::ProviderRetry) -> mentra::runtime::RunOptions { + mentra::runtime::RunOptions::default().with_provider_retry(retry) + } + + fn over(transport: mentra::provider::ResponsesTransport) -> mentra::RuntimeBuilder { + Runtime::empty_builder().with_responses_transport(transport) + } + + let options = patient(mentra::runtime::ProviderRetry { + base_delay: std::time::Duration::from_secs(1), + max_delay: std::time::Duration::from_secs(30), + ..Default::default() + }); + let _ = over(mentra::provider::ResponsesTransport::HttpSse); + + // The delegation contract basis depends on, asserted from outside: a + // subagent meets the same provider with the same patience, so a host + // states its schedule once rather than at every boundary. + let child = options.child(); + assert_eq!(child.provider_retry, options.provider_retry); + assert_eq!(child.retry_budget, options.retry_budget); +} + +/// The downstream shape this exists for, written from outside the crate: a +/// host registers an executor that serves named targets and a tool that names +/// one. Every guard around a shell command still applies — only the executor +/// reads the name. Compiling is half the claim; the other half is that the +/// name survives the trip. +#[tokio::test] +async fn a_tool_can_name_the_executor_a_command_runs_on() { + #[derive(Clone, Default)] + struct TargetLog(Arc>>>); + + #[async_trait] + impl RuntimeExecutor for TargetLog { + async fn run(&self, request: CommandRequest) -> Result { + self.0 + .lock() + .expect("target log poisoned") + .push(request.target.clone()); + Ok(CommandOutput { + stdout: format!("ran on {}", request.target.unwrap_or("local".to_string())), + stderr: String::new(), + success: true, + status_code: Some(0), + timed_out: false, + stdout_truncated: false, + stderr_truncated: false, + }) + } + } + + struct TargetedShellTool; + + #[async_trait] + impl ToolDefinition for TargetedShellTool { + fn descriptor(&self) -> ToolSpec { + ToolSpec::builder("targeted_shell") + .description("Run a command on a named host") + .input_schema(json!({ + "type": "object", + "properties": { + "target": { "type": "string" }, + "command": { "type": "string" } + }, + "required": ["command"] + })) + .build() + } + } + + #[async_trait] + impl ToolExecutor for TargetedShellTool { + async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult { + let command = input + .get("command") + .and_then(Value::as_str) + .ok_or_else(|| "command is required".to_string())? + .to_string(); + let target = input + .get("target") + .and_then(Value::as_str) + .map(str::to_string); + let cwd = ctx.resolve_working_directory(None)?; + let output = ctx + .execute_shell_command_on(target, command, None, None, cwd) + .await?; + Ok(output.stdout) + } + } + + let log = TargetLog::default(); + let model = ModelInfo::new("mock-model", BuiltinProvider::OpenAI); + let provider = ScriptedProvider::new(model.provider.clone(), vec![model.clone()]); + provider.push_turns(vec![ + Turn::ToolCalls(vec![ScriptedToolCall::new( + "targeted_shell", + json!({ "target": "mac", "command": "xcodebuild -version" }), + )]), + Turn::Text("done".to_string()), + ]); + + let runtime = Runtime::builder() + .with_runtime_identifier(format!("public-api-{}", now_nanos())) + .with_store(VolatileRuntimeStore::new()) + .with_provider_instance(provider) + .with_policy(RuntimePolicy::permissive()) + .with_executor(log.clone()) + .build() + .expect("build runtime"); + runtime.register_tool(TargetedShellTool); + let mut agent = runtime.spawn("target-agent", model).expect("spawn agent"); + + agent + .send(vec![ContentBlock::text("build it on the mac")]) + .await + .expect("run completes"); + + assert_eq!( + log.0.lock().expect("target log poisoned").as_slice(), + [Some("mac".to_string())] + ); +} + +fn now_nanos() -> u128 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() +} + +fn spawn_models_server(response_body: &str) -> (String, thread::JoinHandle) { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind test server"); + let addr = listener.local_addr().expect("read listener addr"); + let response_body = response_body.to_string(); + + let handle = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("accept request"); + let mut request = Vec::new(); + let mut temp = [0_u8; 1024]; + + loop { + let read = stream.read(&mut temp).expect("read request"); + if read == 0 { + break; + } + request.extend_from_slice(&temp[..read]); + if request.windows(4).any(|window| window == b"\r\n\r\n") { + break; + } + } + + let response = format!( + concat!( + "HTTP/1.1 200 OK\r\n", + "content-type: application/json\r\n", + "content-length: {}\r\n\r\n", + "{}" + ), + response_body.len(), + response_body + ); + stream + .write_all(response.as_bytes()) + .expect("write response"); + + String::from_utf8(request).expect("request should be valid utf8") + }); + + (format!("http://{addr}/"), handle) +} + +fn model_with_created_at(id: &str, provider: BuiltinProvider, unix_timestamp: i64) -> ModelInfo { + let mut model = ModelInfo::new(id, provider); + model.created_at = Some( + time::OffsetDateTime::from_unix_timestamp(unix_timestamp) + .expect("timestamp should be valid"), + ); + model +} diff --git a/vendor/mentra/tests/responses_transport.rs b/vendor/mentra/tests/responses_transport.rs new file mode 100644 index 0000000..0dee228 --- /dev/null +++ b/vendor/mentra/tests/responses_transport.rs @@ -0,0 +1,254 @@ +//! Which transport a runtime's Responses requests go out on, and what happens +//! when a provider cannot serve the one it was handed. + +use std::{ + fs, + sync::{ + Arc, Mutex, + atomic::{AtomicU64, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use mentra::{ + AgentConfig, BuiltinProvider, ContentBlock, ProviderCapabilities, Role, Runtime, + provider::{ + ContentBlockDelta, ContentBlockStart, ModelInfo, Provider, ProviderDescriptor, + ProviderError, ProviderEvent, ProviderEventStream, ProviderId, ProviderRequestOptions, + Request, ResponsesTransport, + }, + runtime::SqliteRuntimeStore, +}; +use tokio::sync::mpsc; + +static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1); + +/// Answers one short text turn and remembers the options it was handed. +/// +/// The options are what this file is about: the transport is settled before the +/// request leaves the runtime, so what arrives here is the only evidence of +/// which one was chosen. +#[derive(Clone)] +struct RecordingProvider { + id: ProviderId, + display_name: Option, + capabilities: ProviderCapabilities, + seen: Arc>>, +} + +impl RecordingProvider { + fn new(id: impl Into, supports_websockets: bool) -> Self { + Self { + id: id.into(), + display_name: None, + capabilities: ProviderCapabilities { + supports_streaming: true, + supports_websockets, + ..ProviderCapabilities::default() + }, + seen: Arc::new(Mutex::new(Vec::new())), + } + } + + fn named(mut self, display_name: &str) -> Self { + self.display_name = Some(display_name.to_string()); + self + } + + fn transports(&self) -> Vec { + self.seen + .lock() + .expect("recorded options") + .iter() + .map(|options| options.responses.transport) + .collect() + } +} + +#[async_trait] +impl Provider for RecordingProvider { + fn descriptor(&self) -> ProviderDescriptor { + ProviderDescriptor { + id: self.id.clone(), + display_name: self.display_name.clone(), + description: None, + } + } + + fn capabilities(&self) -> ProviderCapabilities { + self.capabilities + } + + async fn list_models(&self) -> Result, ProviderError> { + Ok(vec![ModelInfo::new("model", self.id.clone())]) + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.seen + .lock() + .expect("recorded options") + .push(request.provider_request_options.clone()); + + let (tx, rx) = mpsc::unbounded_channel(); + for event in [ + ProviderEvent::MessageStarted { + id: "msg-1".to_string(), + model: "model".to_string(), + role: Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("ok".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageStopped, + ] { + tx.send(Ok(event)).expect("test receiver alive"); + } + Ok(rx) + } +} + +fn temp_store() -> SqliteRuntimeStore { + let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = std::env::temp_dir().join(format!( + "mentra-responses-transport-{timestamp}-{unique}.sqlite" + )); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).expect("create temp dir"); + } + SqliteRuntimeStore::new(path) +} + +/// A runtime around `provider`, optionally told which transport to use. +fn runtime_for(provider: RecordingProvider, transport: Option) -> Runtime { + let mut builder = Runtime::empty_builder() + .with_provider_instance(provider) + .with_store(temp_store()); + if let Some(transport) = transport { + builder = builder.with_responses_transport(transport); + } + builder.build().expect("build runtime") +} + +#[tokio::test] +async fn a_runtime_that_chooses_nothing_still_streams_over_http_sse() { + let provider = RecordingProvider::new(BuiltinProvider::OpenAI, true); + let recorder = provider.clone(); + let runtime = runtime_for(provider, None); + + let mut agent = runtime + .spawn("agent", ModelInfo::new("model", BuiltinProvider::OpenAI)) + .expect("spawn agent"); + agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect("the turn runs"); + + assert_eq!(recorder.transports(), vec![ResponsesTransport::HttpSse]); +} + +#[test] +fn a_runtime_reports_the_transport_it_was_given() { + // Without this reader the choice is write-only: the only evidence a host's + // selection reached the runtime is a turn run against a provider that + // records what it was handed, so anything downstream can test its own field + // and stop at the seam — the shape of test that passes while the wiring + // between the two is broken. + let unset = runtime_for(RecordingProvider::new(BuiltinProvider::OpenAI, true), None); + assert_eq!(unset.responses_transport(), None); + + let chosen = runtime_for( + RecordingProvider::new(BuiltinProvider::OpenAI, true), + Some(ResponsesTransport::WebSocket), + ); + assert_eq!( + chosen.responses_transport(), + Some(ResponsesTransport::WebSocket) + ); +} + +#[tokio::test] +async fn a_chosen_websocket_transport_reaches_the_request() { + // The gap this closes: `ResponsesRequestOptions.transport` existed and the + // websocket path was compiled in, but nothing in the runtime ever set the + // field, so every request went out over HTTP+SSE whatever the host wanted. + let provider = RecordingProvider::new(BuiltinProvider::OpenAI, true); + let recorder = provider.clone(); + let runtime = runtime_for(provider, Some(ResponsesTransport::WebSocket)); + + let mut agent = runtime + .spawn("agent", ModelInfo::new("model", BuiltinProvider::OpenAI)) + .expect("spawn agent"); + agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect("the turn runs"); + + assert_eq!(recorder.transports(), vec![ResponsesTransport::WebSocket]); +} + +#[tokio::test] +async fn a_runtime_choice_settles_a_disagreeing_agent_config() { + // Two live opinions about one socket is not a state the runtime keeps: the + // connection-level answer is the one that holds. + let provider = RecordingProvider::new(BuiltinProvider::OpenAI, true); + let recorder = provider.clone(); + let runtime = runtime_for(provider, Some(ResponsesTransport::HttpSse)); + + let mut config = AgentConfig::default(); + config.provider_request_options.responses.transport = ResponsesTransport::WebSocket; + let mut agent = runtime + .spawn_with_config( + "agent", + ModelInfo::new("model", BuiltinProvider::OpenAI), + config, + ) + .expect("spawn agent"); + agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect("the turn runs"); + + assert_eq!(recorder.transports(), vec![ResponsesTransport::HttpSse]); +} + +#[tokio::test] +async fn a_provider_without_websockets_refuses_rather_than_pretending() { + // anthropic and gemini report `supports_websockets: false`. Answering over + // HTTP+SSE would look like success and be a transport nobody asked for. + let provider = RecordingProvider::new(BuiltinProvider::Anthropic, false).named("Anthropic"); + let recorder = provider.clone(); + let runtime = runtime_for(provider, Some(ResponsesTransport::WebSocket)); + + let mut agent = runtime + .spawn("agent", ModelInfo::new("model", BuiltinProvider::Anthropic)) + .expect("spawn agent"); + let error = agent + .send(vec![ContentBlock::text("hello")]) + .await + .expect_err("a transport the provider cannot serve is refused"); + + let message = error.to_string(); + assert!( + message.contains("Anthropic"), + "the refusal must name the provider: {message}" + ); + assert!( + message.contains("websocket"), + "and say what it could not do: {message}" + ); + assert!( + recorder.transports().is_empty(), + "the request must never have been sent" + ); +} diff --git a/vendor/mentra/tests/skills_api.rs b/vendor/mentra/tests/skills_api.rs new file mode 100644 index 0000000..7be9ce9 --- /dev/null +++ b/vendor/mentra/tests/skills_api.rs @@ -0,0 +1,145 @@ +//! Public-API tests for skill registration and enumeration. +//! +//! These exercise what a host can actually reach: registering roots through +//! `Runtime`, listing what loaded, and naming the error type in its own +//! signatures. + +use std::{ + fs, + path::{Path, PathBuf}, + sync::atomic::{AtomicU64, Ordering}, + time::{SystemTime, UNIX_EPOCH}, +}; + +use mentra::{BuiltinProvider, Runtime, SkillInfo, SkillLoadError}; + +static NEXT_ID: AtomicU64 = AtomicU64::new(0); + +fn temp_dir(label: &str) -> PathBuf { + let unique = NEXT_ID.fetch_add(1, Ordering::Relaxed); + let stamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system time") + .as_nanos(); + let path = std::env::temp_dir().join(format!("mentra-skills-api-{label}-{stamp}-{unique}")); + fs::create_dir_all(&path).expect("create temp dir"); + path +} + +fn write_skill(root: &Path, dir: &str, name: &str, description: &str, body: &str) { + let skill_dir = root.join(dir); + fs::create_dir_all(&skill_dir).expect("create skill dir"); + fs::write( + skill_dir.join("SKILL.md"), + format!("---\nname: {name}\ndescription: {description}\n---\n{body}\n"), + ) + .expect("write skill"); +} + +fn runtime() -> Runtime { + Runtime::builder() + .with_provider(BuiltinProvider::OpenAI, "test-key") + .build() + .expect("runtime builds") +} + +/// The error type must be nameable by a caller — this signature is the test. +fn register(runtime: &Runtime, path: &Path) -> Result<(), SkillLoadError> { + runtime.register_skills_dir(path) +} + +#[test] +fn a_host_can_name_the_error_type() { + let runtime = runtime(); + let missing = temp_dir("nameable").join("does-not-exist"); + + let error = register(&runtime, &missing).expect_err("an unreadable root is an error"); + + // And match on it, which is the point of it being an enum. + assert!(matches!(error, SkillLoadError::ReadDir { .. })); +} + +#[test] +fn registering_a_second_root_keeps_the_first() { + let workspace = temp_dir("layer-workspace"); + write_skill(&workspace, "review", "review", "project review", "PROJECT"); + let global = temp_dir("layer-global"); + write_skill(&global, "review", "review", "personal review", "PERSONAL"); + write_skill(&global, "deploy", "deploy", "personal deploy", "DEPLOY"); + + let runtime = runtime(); + runtime + .register_skills_dirs([workspace.as_path(), global.as_path()]) + .expect("both roots register"); + + let skills = runtime.skills(); + let names: Vec<&str> = skills.iter().map(|skill| skill.name.as_str()).collect(); + assert_eq!(names, vec!["deploy", "review"]); + + let review = skills + .iter() + .find(|skill| skill.name == "review") + .expect("review present"); + assert_eq!( + review.description, "project review", + "the earlier root must win the collision" + ); + assert!(review.path.starts_with(&workspace)); +} + +#[test] +fn enumeration_reports_name_description_and_source() { + let root = temp_dir("enumerate"); + write_skill(&root, "haiku", "haiku", "writes haiku", "BODY"); + + let runtime = runtime(); + runtime.register_skills_dir(&root).expect("registers"); + + let skills = runtime.skills(); + assert_eq!(skills.len(), 1); + let SkillInfo { + name, + description, + path, + } = &skills[0]; + assert_eq!(name, "haiku"); + assert_eq!(description, "writes haiku"); + assert_eq!(path, &root.join("haiku").join("SKILL.md")); +} + +#[test] +fn a_runtime_without_skills_lists_none() { + assert!(runtime().skills().is_empty()); +} + +#[test] +fn a_duplicate_name_inside_one_root_is_still_an_error() { + let root = temp_dir("duplicate"); + write_skill(&root, "first", "shared", "one", "A"); + write_skill(&root, "second", "shared", "two", "B"); + + let error = runtime() + .register_skills_dir(&root) + .expect_err("a repeated name in one root is a mistake"); + + assert!(matches!(error, SkillLoadError::DuplicateSkillName { .. })); +} + +#[test] +fn roots_before_a_failing_one_stay_registered() { + let good = temp_dir("partial-good"); + write_skill(&good, "keep", "keep", "kept", "BODY"); + let missing = temp_dir("partial-missing").join("absent"); + + let runtime = runtime(); + let error = runtime + .register_skills_dirs([good.as_path(), missing.as_path()]) + .expect_err("the second root fails"); + + assert!(matches!(error, SkillLoadError::ReadDir { .. })); + assert_eq!( + runtime.skills().len(), + 1, + "the root that loaded before the failure stays registered" + ); +}