diff --git a/Cargo.lock b/Cargo.lock index aadf0f7..e6434d1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -817,6 +817,7 @@ dependencies = [ "headless_chrome", "iced", "image", + "libc", "memmap2", "muda", "png 0.17.16", @@ -1198,6 +1199,7 @@ dependencies = [ "libc", "libgit2-sys", "log", + "openssl-probe", "openssl-sys", "url", ] @@ -2558,6 +2560,12 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +[[package]] +name = "openssl-probe" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d05e27ee213611ffe7d6348b942e8f942b37114c00cc03cec254295a4a17852e" + [[package]] name = "openssl-src" version = "300.6.1+3.6.3" diff --git a/Cargo.toml b/Cargo.toml index 962d987..9af7a3f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -14,10 +14,11 @@ cc = "1.3.0" [dependencies] diesel = { version = "2.3.11", features = ["sqlite", "returning_clauses_for_sqlite_3_35", "64-column-tables"] } diesel_migrations = "2.3.2" -git2 = { version = "0.21.0", features = ["cred", "vendored-libgit2", "vendored-openssl"] } +git2 = { version = "0.21.0", features = ["https", "vendored-libgit2", "vendored-openssl"] } headless_chrome = "1.0.22" iced = { version = "0.14.0", default-features = false, features = ["advanced", "image-without-codecs", "markdown", "svg", "tokio", "wgpu"] } image = { version = "0.25.10", default-features = false, features = ["gif", "jpeg", "png", "webp"] } +libc = "0.2.186" memmap2 = "0.9.11" png = "0.17.16" regex = "1.13.1" diff --git a/src/agent.rs b/src/agent.rs index 233946b..8e64ee5 100644 --- a/src/agent.rs +++ b/src/agent.rs @@ -40,10 +40,10 @@ struct AgentSkillMetadata { description: String, } -struct AgentSkill { - name: String, - description: String, - path: PathBuf, +pub(crate) struct AgentSkill { + pub(crate) name: String, + pub(crate) description: String, + pub(crate) path: PathBuf, } fn agent_skills_root(home: Option<&Path>) -> Option { @@ -76,7 +76,7 @@ fn skill_frontmatter(content: &str) -> Option<&str> { Some(content[..end].trim_end_matches('\r')) } -fn discover_agent_skills(root: &Path) -> Vec { +pub(crate) fn discover_agent_skills(root: &Path) -> Vec { let Ok(root) = root.canonicalize() else { return Vec::new(); }; @@ -109,8 +109,17 @@ fn discover_agent_skills(root: &Path) -> Vec { skills } +#[cfg(test)] fn agent_skills_prompt_for(root: &Path) -> Option { - let skills = discover_agent_skills(root); + agent_skills_prompt_for_roots(std::iter::once(root)) +} + +fn agent_skills_prompt_for_roots<'a>(roots: impl IntoIterator) -> Option { + let mut skills = roots + .into_iter() + .flat_map(discover_agent_skills) + .collect::>(); + skills.sort_by(|left, right| left.name.cmp(&right.name)); if skills.is_empty() { return None; } @@ -131,9 +140,14 @@ fn agent_skills_prompt_for(root: &Path) -> Option { )) } -pub(crate) fn agent_skills_prompt() -> Option { - let root = agent_skills_root(std::env::var_os("HOME").as_deref().map(Path::new))?; - agent_skills_prompt_for(&root) +pub(crate) fn agent_skills_prompt_with(extension_roots: &[PathBuf]) -> Option { + let standard = agent_skills_root(std::env::var_os("HOME").as_deref().map(Path::new)); + agent_skills_prompt_for_roots( + standard + .iter() + .map(PathBuf::as_path) + .chain(extension_roots.iter().map(PathBuf::as_path)), + ) } fn user_shell() -> OsString { @@ -1526,6 +1540,8 @@ struct RalphRuntime { engine: EngineSettings, turn: TurnSettings, idle_timeout: Duration, + extensions: crate::extensions::ExtensionRegistry, + session_id: i32, } #[cfg(target_os = "macos")] @@ -1612,6 +1628,17 @@ impl Tools { Ok(()) } + pub(crate) fn enable_agent_skill_roots(&mut self, roots: &[PathBuf]) { + self.agent_skill_roots.extend( + roots + .iter() + .flat_map(|root| discover_agent_skills(root)) + .filter_map(|skill| skill.path.parent().map(Path::to_owned)), + ); + self.agent_skill_roots.sort(); + self.agent_skill_roots.dedup(); + } + #[cfg(target_os = "macos")] pub(crate) fn enable_ralph( &mut self, @@ -1620,13 +1647,17 @@ impl Tools { engine: EngineSettings, turn: TurnSettings, idle_timeout: Duration, + extension_context: (crate::extensions::ExtensionRegistry, i32), ) { + let (extensions, session_id) = extension_context; self.ralph = Some(RalphRuntime { service, model, engine, turn, idle_timeout, + extensions, + session_id, }); } @@ -1786,6 +1817,53 @@ impl Tools { // A unique cache namespace preserves tool continuations inside this // round while preventing parent or earlier-round KV restoration. let directory = RalphRoundDirectory::create(round)?; + let mut round_turn = runtime.turn.clone(); + let hooks = runtime.extensions.dispatch( + &[crate::extensions::HookEvent::SubagentStart { + session_id: runtime.session_id, + round, + }], + &self.root, + &runtime.model.to_string(), + cancel, + ); + let _ = runtime.extensions.persist_hook_results(&hooks); + if let Some(status) = hooks.status() { + send_state( + events, + index, + ToolLifecycle::Running, + Some(format!("Ralph round {round}/{max_rounds} · {status}\n")), + ); + } + if !hooks.errors.is_empty() { + send_state( + events, + index, + ToolLifecycle::Running, + Some(format!( + "Ralph round {round}/{max_rounds} · extension hook warning: {}\n", + hooks + .errors + .iter() + .map(|(id, error)| format!("{id}: {error}")) + .collect::>() + .join("; ") + )), + ); + } + for output in &hooks.outputs { + if let Some(context) = &output.additional_context { + round_turn.system_prompt.push_str("\n\n"); + round_turn + .system_prompt + .push_str(&crate::extensions::wrap_context( + &output.extension_id, + &output.event, + context, + )); + } + } let mut messages = vec![ChatTurn { user: true, tool: false, @@ -1804,7 +1882,8 @@ impl Tools { "Ralph round {round}/{max_rounds} · child generation {step}\n" )), ); - let output = run_ralph_generation(runtime, &messages, &directory.0, cancel)?; + let output = + run_ralph_generation(runtime, &round_turn, &messages, &directory.0, cancel)?; let calls = parse_tool_calls(runtime.model, &output.message.content) .map_err(|error| format!("malformed child tool syntax: {error}"))? .1; @@ -2622,13 +2701,14 @@ impl Tools { #[cfg(target_os = "macos")] fn run_ralph_generation( runtime: &RalphRuntime, + turn: &TurnSettings, messages: &[ChatTurn], checkpoint: &Path, cancel: &AtomicBool, ) -> Result { let active = runtime.service.generate( runtime.engine.clone(), - runtime.turn.clone(), + turn.clone(), messages.to_vec(), CheckpointTarget::Transient(checkpoint.to_owned()), WorkSource::LocalChat, diff --git a/src/app.rs b/src/app.rs index a38e835..c41a14c 100644 --- a/src/app.rs +++ b/src/app.rs @@ -1,3 +1,4 @@ +mod extensions; mod generation; mod git; mod model_manager; @@ -37,7 +38,7 @@ use crate::settings::{ use iced::widget::{markdown, scrollable, text_editor}; use iced::{Size, Subscription, Task, keyboard, mouse, window}; use rfd::AsyncFileDialog; -use std::collections::{BTreeMap, HashMap, HashSet, VecDeque}; +use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet, VecDeque}; use std::fs; use std::path::{Path, PathBuf}; use std::sync::Arc; @@ -73,6 +74,13 @@ pub(crate) struct App { config: Config, preference_draft: PreferenceDraft, preference_error: Option, + pub(super) extensions: crate::extensions::ExtensionRegistry, + pub(super) extension_source: String, + pub(super) extension_ref: String, + pub(super) extension_error: Option, + extension_operation: Option, + pub(super) pending_extension_trust: Option, + pub(super) pending_extension_uninstall: Option, restore_dev_brain_confirmation: bool, selected_project: Option, selected_session: Option, @@ -150,6 +158,8 @@ pub(crate) struct App { #[cfg(target_os = "macos")] active_tool_check: Option, #[cfg(target_os = "macos")] + active_extension_hooks: Option, + #[cfg(target_os = "macos")] agent_tools: Option<(i32, Arc>)>, #[cfg(target_os = "macos")] active_tools: Option, @@ -206,6 +216,7 @@ struct ChatSnapshot { active_generation: Option, active_compaction: Option, active_tool_check: Option, + active_extension_hooks: Option, agent_tools: Option<(i32, Arc>)>, active_tools: Option, tool_cards: Vec, @@ -260,6 +271,7 @@ pub(super) enum PreferenceSection { Model, Endpoint, DevBrain, + Extensions, Git, Prompt, Generation, @@ -270,10 +282,11 @@ pub(super) enum PreferenceSection { } impl PreferenceSection { - const ALL: [Self; 10] = [ + const ALL: [Self; 11] = [ Self::Model, Self::Endpoint, Self::DevBrain, + Self::Extensions, Self::Git, Self::Prompt, Self::Generation, @@ -288,6 +301,7 @@ impl PreferenceSection { Self::Model => "preferences-model", Self::Endpoint => "preferences-endpoint", Self::DevBrain => "preferences-dev-brain", + Self::Extensions => "preferences-extensions", Self::Git => "preferences-git", Self::Prompt => "preferences-prompt", Self::Generation => "preferences-generation", @@ -303,6 +317,7 @@ impl PreferenceSection { Self::Model => "Model & lifecycle", Self::Endpoint => "Local endpoint", Self::DevBrain => "Dev Brain", + Self::Extensions => "Agent extensions", Self::Git => "Git diffs", Self::Prompt => "Prompt", Self::Generation => "Generation", @@ -368,6 +383,17 @@ pub(crate) enum Message { PreferenceEndpointCorsChanged(bool), PreferenceDevBrainEnabledChanged(bool), PreferenceDevBrainVaultChanged(String), + PreferenceExtensionSourceChanged(String), + PreferenceExtensionRefChanged(String), + InstallExtension, + UpdateExtension(String), + ToggleExtension(String, bool), + ConfirmExtensionTrust, + CancelExtensionTrust, + RequestUninstallExtension(String), + ConfirmUninstallExtension, + CancelUninstallExtension, + ExtensionOperationTick, PreferenceGitDiffLayoutChanged(GitDiffLayout), PreferenceGitDiffAlgorithmChanged(GitDiffAlgorithm), PreferenceGitContextLinesChanged(String), @@ -497,6 +523,11 @@ pub(crate) enum Message { ClearTransientCache, } +pub(super) struct ActiveExtensionOperation { + label: String, + receiver: mpsc::Receiver>, +} + impl App { pub(crate) fn load(main_window: window::Id) -> Self { let config = match Config::load(&config_path()) { @@ -528,6 +559,15 @@ impl App { Arc::new(Metrics::new(&application_support_path().join("kv-cache"))); let metrics_snapshot = metrics.snapshot(); let git_diff_layout = config.git.diff_layout; + let extension_root = extensions_path(); + let (extensions, extension_error) = + match crate::extensions::ExtensionRegistry::load(&extension_root) { + Ok(extensions) => (extensions, None), + Err(error) => ( + crate::extensions::ExtensionRegistry::empty(&extension_root), + Some(error), + ), + }; #[cfg(target_os = "macos")] let (runtime_config, generation_service, endpoint, service_error) = spawn_services(&config, Arc::clone(&metrics)); @@ -548,6 +588,13 @@ impl App { config, preference_draft, preference_error: None, + extensions, + extension_source: String::new(), + extension_ref: String::new(), + extension_error, + extension_operation: None, + pending_extension_trust: None, + pending_extension_uninstall: None, restore_dev_brain_confirmation: false, selected_project: last_project, selected_session: None, @@ -617,6 +664,7 @@ impl App { active_compaction: None, #[cfg(target_os = "macos")] active_tool_check: None, + active_extension_hooks: None, #[cfg(target_os = "macos")] agent_tools: None, #[cfg(target_os = "macos")] @@ -671,6 +719,15 @@ impl App { let metrics = Arc::new(Metrics::new(&application_support_path().join("kv-cache"))); let metrics_snapshot = metrics.snapshot(); let git_diff_layout = config.git.diff_layout; + let extension_root = extensions_path(); + let (extensions, extension_error) = + match crate::extensions::ExtensionRegistry::load(&extension_root) { + Ok(extensions) => (extensions, None), + Err(error) => ( + crate::extensions::ExtensionRegistry::empty(&extension_root), + Some(error), + ), + }; #[cfg(target_os = "macos")] let (runtime_config, generation_service, endpoint, service_error) = spawn_services(&config, Arc::clone(&metrics)); @@ -695,6 +752,13 @@ impl App { config, preference_draft, preference_error: None, + extensions, + extension_source: String::new(), + extension_ref: String::new(), + extension_error, + extension_operation: None, + pending_extension_trust: None, + pending_extension_uninstall: None, restore_dev_brain_confirmation: false, selected_project: None, selected_session: None, @@ -762,6 +826,7 @@ impl App { active_compaction: None, #[cfg(target_os = "macos")] active_tool_check: None, + active_extension_hooks: None, #[cfg(target_os = "macos")] agent_tools: None, #[cfg(target_os = "macos")] @@ -826,6 +891,7 @@ impl App { active_generation: self.active_generation.take(), active_compaction: self.active_compaction.take(), active_tool_check: self.active_tool_check.take(), + active_extension_hooks: self.active_extension_hooks.take(), agent_tools: self.agent_tools.take(), active_tools: self.active_tools.take(), tool_cards: std::mem::take(&mut self.tool_cards), @@ -872,6 +938,7 @@ impl App { self.active_generation = snapshot.active_generation; self.active_compaction = snapshot.active_compaction; self.active_tool_check = snapshot.active_tool_check; + self.active_extension_hooks = snapshot.active_extension_hooks; self.agent_tools = snapshot.agent_tools; self.active_tools = snapshot.active_tools; self.tool_cards = snapshot.tool_cards; @@ -991,6 +1058,8 @@ impl App { if let Some(id) = self.preferences_window { self.preference_error = None; self.restore_dev_brain_confirmation = false; + self.pending_extension_trust = None; + self.pending_extension_uninstall = None; return window::close(id); } } @@ -1083,6 +1152,8 @@ impl App { self.preferences_window = None; self.preference_error = None; self.restore_dev_brain_confirmation = false; + self.pending_extension_trust = None; + self.pending_extension_uninstall = None; } if self.help_window == Some(id) { self.help_window = None; @@ -1411,6 +1482,10 @@ impl App { if let Some(check) = &self.active_tool_check { check.active.cancel.store(true, Ordering::Relaxed); } + #[cfg(target_os = "macos")] + if let Some(hooks) = &self.active_extension_hooks { + hooks.cancel.store(true, Ordering::Relaxed); + } } Message::GenerationTick => { #[cfg(target_os = "macos")] @@ -1837,6 +1912,8 @@ impl App { self.permission_mode = permission_mode; self.error = restore_error; self.reload_projects(); + #[cfg(target_os = "macos")] + self.start_resume_extension_hooks(); return Task::batch([scroll_chat_to_end(), self.load_next_a2ui_image()]); } Err(error) => { @@ -1977,6 +2054,12 @@ impl App { iced::time::every(Duration::from_millis(100)).map(|_| Message::GitOperationTick), ); } + if self.extension_operation.is_some() { + subscriptions.push( + iced::time::every(Duration::from_millis(100)) + .map(|_| Message::ExtensionOperationTick), + ); + } #[cfg(target_os = "macos")] let titling = self.active_titling.is_some(); #[cfg(not(target_os = "macos"))] @@ -2116,6 +2199,7 @@ impl App { } fn quit(&mut self) -> Task { + self.cancel_running_extension_hooks(); if self.database.is_none() { return iced::exit(); } @@ -2403,6 +2487,10 @@ pub(crate) fn browser_profile_path() -> PathBuf { application_support_path().join("browser") } +pub(crate) fn extensions_path() -> PathBuf { + application_support_path().join("extensions") +} + /// The settings file, beside the project database. pub(crate) fn config_path() -> PathBuf { application_support_path().join("config.yaml") diff --git a/src/app/extensions.rs b/src/app/extensions.rs new file mode 100644 index 0000000..ea43f4d --- /dev/null +++ b/src/app/extensions.rs @@ -0,0 +1,159 @@ +use super::*; + +impl App { + pub(super) fn install_extension(&mut self) { + let source = self.extension_source.trim().to_owned(); + let requested_ref = self.extension_ref.trim().to_owned(); + if source.is_empty() { + self.extension_error = Some("Enter an HTTPS Git repository URL.".into()); + return; + } + self.cancel_running_extension_hooks(); + let root = extensions_path(); + self.start_extension_operation("Installing extension", move || { + crate::extensions::ExtensionRegistry::install( + &root, + &source, + (!requested_ref.is_empty()).then_some(requested_ref.as_str()), + ) + }); + } + + pub(super) fn update_extension(&mut self, id: String) { + self.cancel_running_extension_hooks(); + let root = extensions_path(); + self.start_extension_operation("Updating extension", move || { + crate::extensions::ExtensionRegistry::update(&root, &id) + }); + } + + pub(super) fn toggle_extension(&mut self, id: String, enabled: bool) { + let Some(extension) = self + .extensions + .extensions + .iter() + .find(|extension| extension.id == id) + else { + return; + }; + if enabled && extension.has_commands() && !extension.trusted { + self.pending_extension_trust = Some(id); + return; + } + self.cancel_running_extension_hooks(); + let root = extensions_path(); + self.start_extension_operation( + if enabled { + "Enabling extension" + } else { + "Disabling extension" + }, + move || crate::extensions::ExtensionRegistry::set_enabled(&root, &id, enabled), + ); + } + + pub(super) fn confirm_extension_trust(&mut self) { + let Some(id) = self.pending_extension_trust.take() else { + return; + }; + self.cancel_running_extension_hooks(); + let root = extensions_path(); + self.start_extension_operation("Enabling trusted extension", move || { + crate::extensions::ExtensionRegistry::trust_and_enable(&root, &id) + }); + } + + pub(super) fn confirm_uninstall_extension(&mut self) { + let Some(id) = self.pending_extension_uninstall.take() else { + return; + }; + self.cancel_running_extension_hooks(); + let root = extensions_path(); + self.start_extension_operation("Uninstalling extension", move || { + crate::extensions::ExtensionRegistry::uninstall(&root, &id) + }); + } + + fn start_extension_operation( + &mut self, + label: &str, + operation: impl FnOnce() -> Result + + Send + + 'static, + ) { + if self.extension_operation.is_some() { + self.extension_error = Some("Another extension operation is already running.".into()); + return; + } + let (sender, receiver) = mpsc::channel(); + thread::spawn(move || { + let _ = sender.send(operation()); + }); + self.extension_operation = Some(ActiveExtensionOperation { + label: label.into(), + receiver, + }); + self.extension_error = None; + } + + pub(super) fn poll_extension_operation(&mut self) { + let Some(operation) = &self.extension_operation else { + return; + }; + match operation.receiver.try_recv() { + Ok(Ok(extensions)) => { + self.extensions = extensions; + self.extension_operation = None; + self.extension_source.clear(); + self.extension_ref.clear(); + self.extension_error = None; + #[cfg(target_os = "macos")] + { + self.agent_tools = None; + } + } + Ok(Err(error)) => { + self.extension_operation = None; + self.extension_error = Some(error); + } + Err(TryRecvError::Empty) => {} + Err(TryRecvError::Disconnected) => { + self.extension_operation = None; + self.extension_error = Some("The extension operation stopped unexpectedly.".into()); + } + } + } + + pub(super) fn extension_operation_label(&self) -> Option<&str> { + self.extension_operation + .as_ref() + .map(|operation| operation.label.as_str()) + } + + pub(super) fn cancel_running_extension_hooks(&mut self) { + #[cfg(target_os = "macos")] + { + let requests = self + .active_extension_hooks + .iter() + .chain( + self.background_chats + .values() + .filter_map(|chat| chat.active_extension_hooks.as_ref()), + ) + .map(|request| (Arc::clone(&request.cancel), Arc::clone(&request.finished))) + .collect::>(); + for (cancel, _) in &requests { + cancel.store(true, Ordering::Relaxed); + } + let deadline = Instant::now() + Duration::from_secs(1); + while requests + .iter() + .any(|(_, finished)| !finished.load(Ordering::Acquire)) + && Instant::now() < deadline + { + thread::sleep(Duration::from_millis(10)); + } + } + } +} diff --git a/src/app/generation.rs b/src/app/generation.rs index e310e34..0f1d92e 100644 --- a/src/app/generation.rs +++ b/src/app/generation.rs @@ -27,6 +27,21 @@ pub(super) enum PendingContinuation { DurableTool(Vec), } +#[cfg(target_os = "macos")] +pub(super) enum HookContinuation { + User { prompt: String, opening_turn: bool }, + Resume, + Compaction(PendingContinuation), +} + +#[cfg(target_os = "macos")] +pub(super) struct ExtensionHookRequest { + pub(super) cancel: Arc, + pub(super) finished: Arc, + receiver: mpsc::Receiver>, + continuation: HookContinuation, +} + #[cfg(target_os = "macos")] pub(super) struct CompactionRequest { pub(super) active: ActiveGeneration, @@ -369,6 +384,34 @@ fn has_chat_after_last_compaction(messages: &[ChatMessage]) -> bool { }) } +fn extension_context_visibility(messages: &[ChatMessage], enabled: &BTreeSet) -> Vec { + let mut latest_session_context = BTreeMap::new(); + for (index, message) in messages.iter().enumerate() { + if let Some((id, "SessionStart")) = message + .system + .then(|| crate::extensions::context_identity(&message.content)) + .flatten() + && enabled.contains(id) + { + latest_session_context.insert(id.to_owned(), index); + } + } + messages + .iter() + .enumerate() + .map(|(index, message)| { + !message.system + || crate::extensions::context_identity(&message.content).is_none_or( + |(id, event)| { + enabled.contains(id) + && (event != "SessionStart" + || latest_session_context.get(id) == Some(&index)) + }, + ) + }) + .collect() +} + #[cfg(any(target_os = "macos", test))] fn title_context(messages: impl IntoIterator) -> Vec { let mut started = false; @@ -384,13 +427,24 @@ fn title_context(messages: impl IntoIterator) -> Vec } impl App { + fn agent_skills_prompt(&self) -> Option { + let roots = self + .extensions + .enabled_skill_roots() + .unwrap_or_default() + .into_iter() + .map(|(_, root)| root) + .collect::>(); + crate::agent::agent_skills_prompt_with(&roots) + } + fn chat_system_prompt(&self, model: ModelChoice, prompt: &str) -> String { let mut prompt = crate::agent::system_prompt(model, prompt, self.config.dev_brain.enabled); if self.config.a2ui_enabled { prompt.push_str("\n\n"); prompt.push_str(crate::a2ui::SYSTEM_PROMPT); } - if let Some(skills) = crate::agent::agent_skills_prompt() { + if let Some(skills) = self.agent_skills_prompt() { prompt.push_str("\n\n"); prompt.push_str(&skills); } @@ -516,6 +570,74 @@ impl App { } return; } + #[cfg(target_os = "macos")] + { + let preview_session_id = self.selected_session.unwrap_or_default(); + let mut events = Vec::with_capacity(2); + if opening_turn { + events.push(crate::extensions::HookEvent::SessionStart { + session_id: preview_session_id, + reason: "startup", + }); + } + events.push(crate::extensions::HookEvent::UserPromptSubmit { + session_id: preview_session_id, + prompt: prompt.clone(), + }); + if self.extensions.has_hooks_for(&events) { + self.composer = text_editor::Content::new(); + let session_id = match self.selected_session { + Some(session_id) => session_id, + None => { + let Some(project_id) = self.selected_project else { + return; + }; + match self.persist_session(project_id) { + Ok(session_id) => session_id, + Err(error) => { + self.error = Some(format!("Could not create the session: {error}")); + return; + } + } + } + }; + for event in &mut events { + match event { + crate::extensions::HookEvent::SessionStart { + session_id: event_session, + .. + } + | crate::extensions::HookEvent::UserPromptSubmit { + session_id: event_session, + .. + } + | crate::extensions::HookEvent::SubagentStart { + session_id: event_session, + .. + } => *event_session = session_id, + } + } + self.start_extension_hooks( + events, + HookContinuation::User { + prompt, + opening_turn, + }, + ); + } else { + self.start_generation_after_hooks(prompt, opening_turn, Vec::new()); + } + } + #[cfg(not(target_os = "macos"))] + self.start_generation_after_hooks(prompt, false, Vec::new()); + } + + fn start_generation_after_hooks( + &mut self, + prompt: String, + opening_turn: bool, + hook_contexts: Vec, + ) { let model = self.config.model; let generation = self.config.active_generation(); let runtime = self.config.runtime_for(model); @@ -553,6 +675,8 @@ impl App { injected_system.push(SystemMessage::plain(crate::agent::datetime_context())); } #[cfg(target_os = "macos")] + injected_system.extend(hook_contexts); + #[cfg(target_os = "macos")] let reminder_injected = self.system_prompt_reminder_due(); #[cfg(target_os = "macos")] if reminder_injected { @@ -596,7 +720,6 @@ impl App { #[cfg(target_os = "macos")] { - // A draft session only reaches the database once there is a turn to store. let session_id = match self.selected_session { Some(session_id) => session_id, None => { @@ -707,6 +830,211 @@ impl App { } } + #[cfg(target_os = "macos")] + fn start_extension_hooks( + &mut self, + events: Vec, + continuation: HookContinuation, + ) { + if !self.extensions.has_hooks_for(&events) { + match continuation { + HookContinuation::User { + prompt, + opening_turn, + } => { + let restore = prompt.clone(); + self.start_generation_after_hooks(prompt, opening_turn, Vec::new()); + if !self.generating && self.composer.text().trim().is_empty() { + self.composer = text_editor::Content::with_text(&restore); + } + } + HookContinuation::Resume => {} + HookContinuation::Compaction(pending) => { + self.finish_compaction_continuation(pending) + } + } + return; + } + let Some(root) = self + .projects + .iter() + .find(|project| Some(project.project.id) == self.selected_project) + .map(|project| PathBuf::from(&project.project.path)) + else { + self.error = Some("The active project is unavailable.".into()); + self.generating = false; + return; + }; + let registry = self.extensions.clone(); + let activity = registry + .status_for(&events) + .unwrap_or_else(|| "Running agent extension hooks…".into()); + let model = self.config.model.to_string(); + let cancel = Arc::new(AtomicBool::new(false)); + let finished = Arc::new(AtomicBool::new(false)); + let worker_cancel = Arc::clone(&cancel); + let worker_finished = Arc::clone(&finished); + let (sender, receiver) = mpsc::channel(); + thread::spawn(move || { + let enabled = registry + .enabled_ids() + .into_iter() + .map(str::to_owned) + .collect::>(); + let result = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + registry.dispatch(&events, &root, &model, &worker_cancel) + })) { + Ok(result) => result, + Err(_) => { + let mut result = crate::extensions::HookBatchResult::default(); + for id in enabled { + result + .errors + .insert(id, "An agent extension hook worker panicked.".into()); + } + result + } + }; + worker_finished.store(true, Ordering::Release); + let _ = sender.send(Ok(result)); + }); + self.active_extension_hooks = Some(ExtensionHookRequest { + cancel, + finished, + receiver, + continuation, + }); + self.generating = true; + self.stop_requested = false; + self.activity = Some(activity); + self.error = None; + } + + #[cfg(target_os = "macos")] + pub(super) fn start_resume_extension_hooks(&mut self) { + let Some(session_id) = self.selected_session else { + return; + }; + self.start_extension_hooks( + vec![crate::extensions::HookEvent::SessionStart { + session_id, + reason: "resume", + }], + HookContinuation::Resume, + ); + } + + #[cfg(target_os = "macos")] + fn poll_extension_hooks(&mut self) -> bool { + let Some(request) = &self.active_extension_hooks else { + return false; + }; + let mut worker_error = None; + let result = match request.receiver.try_recv() { + Ok(Ok(result)) => result, + Ok(Err(error)) => { + worker_error = Some(error); + crate::extensions::HookBatchResult::default() + } + Err(TryRecvError::Empty) => return false, + Err(TryRecvError::Disconnected) => { + worker_error = Some("An agent extension hook stopped unexpectedly.".into()); + crate::extensions::HookBatchResult::default() + } + }; + let request = self.active_extension_hooks.take().unwrap(); + if let Err(error) = self.extensions.record_hook_results(&result) { + self.error = Some(format!("Could not save extension hook status: {error}")); + } + let mut notices = Vec::new(); + notices.extend(worker_error); + if let Some(status) = result.status() { + notices.push(status); + } + notices.extend( + result + .errors + .iter() + .map(|(id, error)| format!("{id}: {error}")), + ); + if !notices.is_empty() { + self.context_notice = Some(notices.join(" · ")); + } + if self.stop_requested { + if let HookContinuation::User { prompt, .. } = request.continuation + && self.composer.text().trim().is_empty() + { + self.composer = text_editor::Content::with_text(&prompt); + } + self.generating = false; + self.activity = Some("Stopped".into()); + return false; + } + let contexts = result + .outputs + .iter() + .filter_map(|output| { + output.additional_context.as_deref().map(|context| { + SystemMessage::plain(crate::extensions::wrap_context( + &output.extension_id, + &output.event, + context, + )) + }) + }) + .collect::>(); + match request.continuation { + HookContinuation::User { + prompt, + opening_turn, + } => { + self.generating = false; + self.activity = None; + let restore = prompt.clone(); + self.start_generation_after_hooks(prompt, opening_turn, contexts); + if !self.generating && self.composer.text().trim().is_empty() { + self.composer = text_editor::Content::with_text(&restore); + } + } + HookContinuation::Resume => { + self.generating = false; + self.activity = None; + if let Err(error) = self.persist_extension_contexts(&contexts) { + self.error = Some(error); + } + } + HookContinuation::Compaction(pending) => { + if let Err(error) = self.persist_extension_contexts(&contexts) { + self.generating = false; + self.activity = Some("Failed".into()); + self.error = Some(error); + } else { + self.finish_compaction_continuation(pending); + } + } + } + true + } + + #[cfg(target_os = "macos")] + fn persist_extension_contexts(&mut self, contexts: &[SystemMessage]) -> Result<(), String> { + if contexts.is_empty() { + return Ok(()); + } + let session_id = self + .selected_session + .ok_or_else(|| "The active session is unavailable.".to_owned())?; + let stored = self + .database + .as_mut() + .ok_or_else(|| "The project database is unavailable.".to_owned())? + .record_system_messages(session_id, contexts) + .map_err(|error| format!("Could not save extension context: {error}"))?; + self.conversation + .extend(stored.into_iter().map(ChatMessage::from)); + Ok(()) + } + fn system_prompt_reminder_due(&self) -> bool { crate::agent::prompt_reminder_due(self.context_used, self.system_prompt_seen_at) } @@ -719,7 +1047,7 @@ impl App { if let Some(skills) = self.dev_brain_skills_prompt() { reminders.push(skills); } - reminders.extend(crate::agent::agent_skills_prompt()); + reminders.extend(self.agent_skills_prompt()); if self.config.a2ui_enabled { reminders.push(crate::a2ui::SYSTEM_PROMPT.to_owned()); } @@ -747,6 +1075,10 @@ impl App { } fn poll_generation_step(&mut self) -> bool { + #[cfg(target_os = "macos")] + if self.active_extension_hooks.is_some() { + return self.poll_extension_hooks(); + } #[cfg(target_os = "macos")] if self.active_compaction.is_some() { return self.poll_compaction(); @@ -1261,6 +1593,13 @@ impl App { .collect::>(); tools.enable_dev_brain(&self.config.dev_brain, &projects)?; } + let extension_skill_roots = self + .extensions + .enabled_skill_roots()? + .into_iter() + .map(|(_, root)| root) + .collect::>(); + tools.enable_agent_skill_roots(&extension_skill_roots); self.agent_tools = Some((session_id, Arc::new(Mutex::new(tools)))); } let tools = Arc::clone(&self.agent_tools.as_ref().unwrap().1); @@ -1293,7 +1632,7 @@ impl App { child_turn.system_prompt.push_str("\n\n"); child_turn.system_prompt.push_str(&skills); } - if let Some(skills) = crate::agent::agent_skills_prompt() { + if let Some(skills) = self.agent_skills_prompt() { child_turn.system_prompt.push_str("\n\n"); child_turn.system_prompt.push_str(&skills); } @@ -1306,6 +1645,7 @@ impl App { effective.engine.clone(), child_turn, idle_timeout, + (self.extensions.clone(), session_id), ); let approval_mode = match self.permission_mode { PermissionMode::Heuristic => crate::agent::ShellApprovalMode::Heuristic, @@ -1767,44 +2107,14 @@ impl App { } self.context_used = compacted.context_tokens; self.system_prompt_seen_at = compacted.context_tokens; - self.generating = false; - self.activity = None; - match request.pending { - PendingContinuation::None => self.start_next_queued(), - PendingContinuation::User(prompt) => { - self.composer = text_editor::Content::with_text(&prompt); - self.skip_compaction_once = true; - self.start_generation(); - } - PendingContinuation::Tool(result) => { - if let Err(error) = self.start_tool_result_check( - result, - ToolCheckStage::ResultAfterCompaction, - ) { - self.generating = false; - self.activity = Some("Failed".into()); - self.error = Some(error); - } - } - PendingContinuation::DurableTool(touched_paths) => { - let continuation = self - .persist_workspace_instructions(&touched_paths) - .and_then(|()| { - self.start_tool_result_check( - crate::agent::ToolRunResult { - content: String::new(), - touched_paths, - }, - ToolCheckStage::InstructionsAfterCompaction, - ) - }); - if let Err(error) = continuation { - self.generating = false; - self.activity = Some("Failed".into()); - self.error = Some(error); - } - } - } + let session_id = self.selected_session.unwrap(); + self.start_extension_hooks( + vec![crate::extensions::HookEvent::SessionStart { + session_id, + reason: "compact", + }], + HookContinuation::Compaction(request.pending), + ); return true; } Err(error) => { @@ -1841,6 +2151,47 @@ impl App { } } + #[cfg(target_os = "macos")] + fn finish_compaction_continuation(&mut self, pending: PendingContinuation) { + self.generating = false; + self.activity = None; + match pending { + PendingContinuation::None => self.start_next_queued(), + PendingContinuation::User(prompt) => { + self.composer = text_editor::Content::with_text(&prompt); + self.skip_compaction_once = true; + self.start_generation(); + } + PendingContinuation::Tool(result) => { + if let Err(error) = + self.start_tool_result_check(result, ToolCheckStage::ResultAfterCompaction) + { + self.generating = false; + self.activity = Some("Failed".into()); + self.error = Some(error); + } + } + PendingContinuation::DurableTool(touched_paths) => { + let continuation = self + .persist_workspace_instructions(&touched_paths) + .and_then(|()| { + self.start_tool_result_check( + crate::agent::ToolRunResult { + content: String::new(), + touched_paths, + }, + ToolCheckStage::InstructionsAfterCompaction, + ) + }); + if let Err(error) = continuation { + self.generating = false; + self.activity = Some("Failed".into()); + self.error = Some(error); + } + } + } + } + #[cfg(target_os = "macos")] fn apply_compaction( &mut self, @@ -1924,13 +2275,24 @@ impl App { fn model_chat_messages(&self) -> Vec<&ChatMessage> { let start = compacted_context_start(&self.conversation); - self.conversation[start..] + let visible = &self.conversation[start..]; + let enabled = self + .extensions + .enabled_ids() + .into_iter() + .map(str::to_owned) + .collect(); + let extension_visibility = extension_context_visibility(visible, &enabled); + visible .iter() - .filter(|message| { + .enumerate() + .filter(|(index, message)| { !message.compaction && !(message.system && message.content.starts_with(LEGACY_AGENTS_PREFIX)) && !(message.instruction_metadata.is_some() && message.content.is_empty()) + && extension_visibility[*index] }) + .map(|(_, message)| message) .collect() } @@ -2121,12 +2483,13 @@ pub(super) fn session_title(reply: &str) -> Option { mod tests { use super::{ ChatMessage, TOOL_PROTOCOL_CORRECTION, TurnSummary, chat_turn, compacted_context_start, - correction_already_sent, has_chat_after_last_compaction, has_misplaced_tool_call, - is_empty_response, promote_legacy_turn_summaries, queued_prompt, sync_a2ui_message, - title_context, + correction_already_sent, extension_context_visibility, has_chat_after_last_compaction, + has_misplaced_tool_call, is_empty_response, promote_legacy_turn_summaries, queued_prompt, + sync_a2ui_message, title_context, }; use crate::engine::ChatTurn; use crate::model::ModelChoice; + use std::collections::BTreeSet; fn assistant(reasoning: Option<&str>, content: &str) -> ChatMessage { ChatMessage { @@ -2450,4 +2813,44 @@ mod tests { history.push(message(7, true, false, false, false)); assert!(has_chat_after_last_compaction(&history)); } + + #[test] + fn extension_context_deduplicates_session_start_and_rearms_after_compaction() { + let context = |id: &str, event: &str, value: &str| ChatMessage { + system: true, + content: crate::extensions::wrap_context(id, event, value), + ..assistant(None, "") + }; + let messages = vec![ + context("ponytail", "SessionStart", "startup rules"), + context("ponytail", "UserPromptSubmit", "same-turn rules"), + context("ponytail", "SessionStart", "compact rules"), + context("disabled", "UserPromptSubmit", "must disappear"), + ]; + let enabled = BTreeSet::from(["ponytail".to_owned()]); + assert_eq!( + extension_context_visibility(&messages, &enabled), + [false, true, true, false] + ); + + let off = vec![ + context("ponytail", "SessionStart", "full rules"), + context("ponytail", "SessionStart", ""), + ]; + assert_eq!(extension_context_visibility(&off, &enabled), [false, true]); + let user_prefix = ChatMessage { + user: true, + system: false, + content: crate::extensions::wrap_context("disabled", "SessionStart", "user text"), + ..assistant(None, "") + }; + assert_eq!( + extension_context_visibility(&[user_prefix], &enabled), + [true] + ); + assert_eq!( + extension_context_visibility(&messages, &BTreeSet::new()), + [false, false, false, false] + ); + } } diff --git a/src/app/preferences.rs b/src/app/preferences.rs index 5e39e35..6cff8f6 100644 --- a/src/app/preferences.rs +++ b/src/app/preferences.rs @@ -482,6 +482,8 @@ impl App { self.preference_draft = PreferenceDraft::from_saved(&self.config); self.preference_error = None; self.restore_dev_brain_confirmation = false; + self.pending_extension_trust = None; + self.pending_extension_uninstall = None; let (id, open) = window::open(window::Settings { size: Size::new(920.0, 700.0), min_size: Some(Size::new(720.0, 480.0)), @@ -639,6 +641,23 @@ impl App { message: Message, ) -> Result> { match message { + Message::PreferenceExtensionSourceChanged(value) => { + self.extension_source = value; + self.extension_error = None; + } + Message::PreferenceExtensionRefChanged(value) => { + self.extension_ref = value; + self.extension_error = None; + } + Message::InstallExtension => self.install_extension(), + Message::UpdateExtension(id) => self.update_extension(id), + Message::ToggleExtension(id, enabled) => self.toggle_extension(id, enabled), + Message::ConfirmExtensionTrust => self.confirm_extension_trust(), + Message::CancelExtensionTrust => self.pending_extension_trust = None, + Message::RequestUninstallExtension(id) => self.pending_extension_uninstall = Some(id), + Message::ConfirmUninstallExtension => self.confirm_uninstall_extension(), + Message::CancelUninstallExtension => self.pending_extension_uninstall = None, + Message::ExtensionOperationTick => self.poll_extension_operation(), Message::PreferenceModelChanged(model) => { if let Err(error) = self.preference_draft.store_generation() { self.preference_error = Some(error); diff --git a/src/app/view/preferences.rs b/src/app/view/preferences.rs index 4d19551..12b5e48 100644 --- a/src/app/view/preferences.rs +++ b/src/app/view/preferences.rs @@ -240,6 +240,122 @@ impl App { ] .spacing(10), ); + let mut installed_extensions = column![].spacing(10); + if self.extensions.extensions.is_empty() { + installed_extensions = installed_extensions.push( + text("No agent extensions are installed.") + .size(12) + .color(muted_text()), + ); + } + for extension in &self.extensions.extensions { + let id = extension.id.clone(); + let update_id = id.clone(); + let uninstall_id = id.clone(); + let enabled = extension.enabled; + let toggle_control = toggle(enabled) + .label(extension.name.clone()) + .on_toggle(move |enabled| Message::ToggleExtension(id.clone(), enabled)); + let details = format!( + "{} • {} • {} skills • hooks: {}", + extension.version, + extension.author, + extension.skill_count, + extension.hook_names(), + ); + let source = format!( + "{}{} • commit {}", + extension.source_url, + extension + .requested_ref + .as_ref() + .map_or_else(String::new, |reference| format!(" @ {reference}")), + &extension.resolved_commit[..extension.resolved_commit.len().min(12)], + ); + let mut row_content = column![ + row![ + toggle_control.width(Length::Fill), + action_button("Update").on_press(Message::UpdateExtension(update_id)), + action_button("Uninstall") + .on_press(Message::RequestUninstallExtension(uninstall_id)), + ] + .spacing(8) + .align_y(Alignment::Center), + text(&extension.description).size(12), + text(details).size(12).color(muted_text()), + text(source).size(11).color(muted_text()), + ] + .spacing(5); + if let Some(error) = &extension.last_error { + row_content = row_content.push( + text(format!("Last hook error: {error}")) + .size(12) + .style(iced::widget::text::danger), + ); + } + if self.pending_extension_trust.as_deref() == Some(extension.id.as_str()) { + row_content = row_content.push( + column![ + text("Trust this extension's command hooks?").size(13), + text("Hooks run local programs with your user permissions. DS4Server isolates their environment, limits runtime and output, and never invokes a shell, but the installed code can still read or change files you can access.") + .size(12), + row![ + action_button("Cancel").on_press(Message::CancelExtensionTrust), + action_button("Trust and enable") + .on_press(Message::ConfirmExtensionTrust), + ] + .spacing(8), + ] + .spacing(7), + ); + } + if self.pending_extension_uninstall.as_deref() == Some(extension.id.as_str()) { + row_content = row_content.push( + column![ + text("Remove this extension and its stored session data?").size(13), + row![ + action_button("Cancel").on_press(Message::CancelUninstallExtension), + action_button("Uninstall").on_press(Message::ConfirmUninstallExtension), + ] + .spacing(8), + ] + .spacing(7), + ); + } + installed_extensions = installed_extensions + .push(row_content) + .push(rule::horizontal(1)); + } + let mut extension_content = column![ + text("Install a portable Codex plugin from an HTTPS Git repository. Updates are manual and preserve the selected ref.") + .size(12), + text_input("https://example.com/owner/extension.git", &self.extension_source) + .on_input(Message::PreferenceExtensionSourceChanged) + .padding(9), + row![ + text_input("Optional branch, tag, or commit", &self.extension_ref) + .on_input(Message::PreferenceExtensionRefChanged) + .padding(9) + .width(Length::Fill), + action_button("Install").on_press(Message::InstallExtension), + ] + .spacing(8) + .align_y(Alignment::Center), + installed_extensions, + ] + .spacing(10); + if let Some(label) = self.extension_operation_label() { + extension_content = extension_content.push(text(format!("{label}…")).size(12)); + } + if let Some(error) = &self.extension_error { + extension_content = + extension_content.push(text(error).size(12).style(iced::widget::text::danger)); + } + let extension_group = preference_group( + PreferenceSection::Extensions, + "AGENT EXTENSIONS", + extension_content, + ); let git_group = preference_group( PreferenceSection::Git, "GIT DIFFS", @@ -743,6 +859,7 @@ impl App { model_group, endpoint_group, dev_brain_group, + extension_group, git_group, prompt_group, generation_group, diff --git a/src/extensions.rs b/src/extensions.rs new file mode 100644 index 0000000..35314e7 --- /dev/null +++ b/src/extensions.rs @@ -0,0 +1,1901 @@ +use regex::Regex; +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::collections::{BTreeMap, BTreeSet}; +use std::fs; +use std::io::{Read, Write}; +use std::os::unix::process::CommandExt; +use std::path::{Component, Path, PathBuf}; +use std::process::{Command, Stdio}; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::thread; +use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; + +const REGISTRY_SCHEMA: u32 = 1; +const MANIFEST_PATH: &str = ".codex-plugin/plugin.json"; +const MAX_HOOK_INPUT: usize = 64 * 1024; +const MAX_HOOK_OUTPUT: usize = 64 * 1024; +const MAX_HOOK_TIMEOUT: u64 = 10; +const CONTEXT_PREFIX: &str = "[DS4Server extension "; + +#[derive(Clone, Debug, Deserialize, Serialize)] +pub(crate) struct HookBinding { + pub(crate) event: String, + pub(crate) matcher: Option, + pub(crate) argv: Vec, + pub(crate) timeout_seconds: u64, + pub(crate) status_message: Option, +} + +#[derive(Clone, Debug, Deserialize, Serialize)] +pub(crate) struct InstalledExtension { + pub(crate) id: String, + pub(crate) name: String, + pub(crate) version: String, + pub(crate) description: String, + pub(crate) author: String, + pub(crate) enabled: bool, + pub(crate) trusted: bool, + pub(crate) source_url: String, + pub(crate) requested_ref: Option, + pub(crate) resolved_commit: String, + pub(crate) skills_path: Option, + pub(crate) skill_count: usize, + pub(crate) hooks_path: Option, + pub(crate) hooks: Vec, + pub(crate) last_error: Option, +} + +impl InstalledExtension { + pub(crate) fn hook_names(&self) -> String { + let mut names = self + .hooks + .iter() + .map(|hook| hook.event.as_str()) + .collect::>() + .into_iter() + .collect::>(); + names.sort_unstable(); + if names.is_empty() { + "None".into() + } else { + names.join(", ") + } + } + + pub(crate) fn has_commands(&self) -> bool { + !self.hooks.is_empty() + } +} + +#[derive(Clone, Debug, Deserialize, Serialize)] +pub(crate) struct ExtensionRegistry { + schema_version: u32, + pub(crate) extensions: Vec, + #[serde(skip)] + root: PathBuf, +} + +impl ExtensionRegistry { + pub(crate) fn load(root: &Path) -> Result { + let path = root.join("registry.json"); + let mut registry = match fs::read(&path) { + Ok(bytes) => serde_json::from_slice::(&bytes) + .map_err(|error| format!("Could not read {}: {error}", path.display()))?, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Self { + schema_version: REGISTRY_SCHEMA, + extensions: Vec::new(), + root: root.to_owned(), + }, + Err(error) => return Err(format!("Could not read {}: {error}", path.display())), + }; + if registry.schema_version != REGISTRY_SCHEMA { + return Err(format!( + "Extension registry schema {} is unsupported; expected {REGISTRY_SCHEMA}", + registry.schema_version + )); + } + registry.root = root.to_owned(); + let mut ids = BTreeSet::new(); + for extension in ®istry.extensions { + if !valid_id(&extension.id) || !ids.insert(extension.id.clone()) { + return Err(format!( + "Extension registry contains an invalid or duplicate ID: {}", + extension.id + )); + } + } + Ok(registry) + } + + pub(crate) fn empty(root: &Path) -> Self { + Self { + schema_version: REGISTRY_SCHEMA, + extensions: Vec::new(), + root: root.to_owned(), + } + } + + pub(crate) fn enabled_ids(&self) -> BTreeSet<&str> { + self.extensions + .iter() + .filter(|extension| extension.enabled) + .map(|extension| extension.id.as_str()) + .collect() + } + + pub(crate) fn enabled_skill_roots(&self) -> Result, String> { + self.extensions + .iter() + .filter(|extension| extension.enabled) + .filter_map(|extension| { + extension.skills_path.as_ref().map(|skills| { + let root = self.package_root(&extension.id).join(skills); + resolve_inside(&self.package_root(&extension.id), &root) + .map(|root| (extension.id.clone(), root)) + }) + }) + .collect() + } + + pub(crate) fn install( + root: &Path, + source_url: &str, + requested_ref: Option<&str>, + ) -> Result { + validate_source_url(source_url)?; + let requested_ref = requested_ref + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_owned); + let mut registry = Self::load(root)?; + fs::create_dir_all(root.join("packages")).map_err(|error| error.to_string())?; + let staging = temporary_path(&root.join("packages"), "staging"); + let result = (|| { + let repository = git2::build::RepoBuilder::new() + .clone(source_url, &staging) + .map_err(|error| format!("Could not clone extension: {error}"))?; + let commit = checkout_requested(&repository, requested_ref.as_deref())?; + let mut extension = + validate_package(&staging, source_url, requested_ref.clone(), commit)?; + if registry + .extensions + .iter() + .any(|installed| installed.id == extension.id) + { + return Err(format!("Extension {} is already installed", extension.id)); + } + extension.enabled = false; + extension.trusted = false; + promote_package(root, &extension.id, &staging, None)?; + let package = root.join("packages").join(&extension.id); + registry.extensions.push(extension); + if let Err(error) = registry.save() { + let _ = fs::remove_dir_all(package); + return Err(error); + } + Ok(registry) + })(); + if staging.exists() { + let _ = fs::remove_dir_all(&staging); + } + result + } + + pub(crate) fn update(root: &Path, id: &str) -> Result { + let mut registry = Self::load(root)?; + let index = registry + .extensions + .iter() + .position(|extension| extension.id == id) + .ok_or_else(|| format!("Extension {id} is not installed"))?; + let previous = registry.extensions[index].clone(); + validate_source_url(&previous.source_url)?; + let staging = temporary_path(&root.join("packages"), "staging"); + let result = (|| { + let repository = git2::build::RepoBuilder::new() + .clone(&previous.source_url, &staging) + .map_err(|error| format!("Could not clone extension update: {error}"))?; + let commit = checkout_requested(&repository, previous.requested_ref.as_deref())?; + let mut updated = validate_package( + &staging, + &previous.source_url, + previous.requested_ref.clone(), + commit, + )?; + if updated.id != previous.id { + return Err(format!( + "Extension update changed its ID from {} to {}", + previous.id, updated.id + )); + } + updated.enabled = previous.enabled; + updated.trusted = previous.trusted; + let backup = promote_package(root, id, &staging, Some("previous"))?; + registry.extensions[index] = updated; + if let Err(error) = registry.save() { + rollback_package(root, id, backup.as_deref()); + return Err(error); + } + if let Some(backup) = backup { + let _ = fs::remove_dir_all(backup); + } + Ok(registry) + })(); + if staging.exists() { + let _ = fs::remove_dir_all(&staging); + } + result + } + + pub(crate) fn set_enabled(root: &Path, id: &str, enabled: bool) -> Result { + let mut registry = Self::load(root)?; + let extension = registry + .extensions + .iter_mut() + .find(|extension| extension.id == id) + .ok_or_else(|| format!("Extension {id} is not installed"))?; + if enabled && extension.has_commands() && !extension.trusted { + return Err(format!( + "Extension {id} command hooks have not been trusted" + )); + } + extension.enabled = enabled; + extension.last_error = None; + registry.validate_skill_conflicts()?; + registry.save()?; + Ok(registry) + } + + pub(crate) fn trust_and_enable(root: &Path, id: &str) -> Result { + let mut registry = Self::load(root)?; + let extension = registry + .extensions + .iter_mut() + .find(|extension| extension.id == id) + .ok_or_else(|| format!("Extension {id} is not installed"))?; + extension.trusted = true; + extension.enabled = true; + extension.last_error = None; + registry.validate_skill_conflicts()?; + registry.save()?; + Ok(registry) + } + + pub(crate) fn uninstall(root: &Path, id: &str) -> Result { + let mut registry = Self::load(root)?; + let index = registry + .extensions + .iter() + .position(|extension| extension.id == id) + .ok_or_else(|| format!("Extension {id} is not installed"))?; + let package = root.join("packages").join(id); + let data = root.join("data").join(id); + let package_tomb = temporary_path(root, "uninstall-package"); + let data_tomb = temporary_path(root, "uninstall-data"); + if package.exists() { + fs::rename(&package, &package_tomb).map_err(|error| { + format!("Could not stage {} for removal: {error}", package.display()) + })?; + } + if data.exists() + && let Err(error) = fs::rename(&data, &data_tomb) + { + if package_tomb.exists() { + let _ = fs::rename(&package_tomb, &package); + } + return Err(format!( + "Could not stage {} for removal: {error}", + data.display() + )); + } + registry.extensions.remove(index); + if let Err(error) = registry.save() { + if package_tomb.exists() { + let _ = fs::rename(&package_tomb, &package); + } + if data_tomb.exists() { + let _ = fs::rename(&data_tomb, &data); + } + return Err(error); + } + let _ = fs::remove_dir_all(package_tomb); + let _ = fs::remove_dir_all(data_tomb); + Ok(registry) + } + + pub(crate) fn record_hook_results(&mut self, results: &HookBatchResult) -> Result<(), String> { + let mut current = Self::load(&self.root)?; + current.apply_hook_results(results); + current.save()?; + *self = current; + Ok(()) + } + + pub(crate) fn persist_hook_results(&self, results: &HookBatchResult) -> Result<(), String> { + let mut current = Self::load(&self.root)?; + current.apply_hook_results(results); + current.save() + } + + fn apply_hook_results(&mut self, results: &HookBatchResult) { + for extension in &mut self.extensions { + if let Some(error) = results.errors.get(&extension.id) { + extension.last_error = Some(error.clone()); + } else if results.invoked.contains(&extension.id) { + extension.last_error = None; + } + } + } + + pub(crate) fn dispatch( + &self, + events: &[HookEvent], + project_root: &Path, + model: &str, + cancel: &AtomicBool, + ) -> HookBatchResult { + let mut result = HookBatchResult::default(); + for extension in self.extensions.iter().filter(|extension| extension.enabled) { + if cancel.load(Ordering::Relaxed) { + break; + } + let package_root = self.package_root(&extension.id); + for event in events { + if cancel.load(Ordering::Relaxed) { + break; + } + let bindings = extension + .hooks + .iter() + .filter(|binding| binding.event == event.name()) + .filter(|binding| matcher_applies(binding.matcher.as_deref(), event)); + for binding in bindings { + if cancel.load(Ordering::Relaxed) { + break; + } + result.invoked.insert(extension.id.clone()); + match run_hook( + extension, + binding, + event, + ( + &package_root, + &self.root.join("data").join(&extension.id), + project_root, + ), + model, + cancel, + ) { + Ok(output) => result.outputs.push(output), + Err(error) => { + result.errors.entry(extension.id.clone()).or_insert(error); + break; + } + } + } + if !cancel.load(Ordering::Relaxed) + && extension.id == "ponytail" + && ponytail_mode_command(event) + { + let turned_off = result.outputs.iter().rev().any(|output| { + output.extension_id == extension.id + && output.event == "UserPromptSubmit" + && output.system_message.as_deref() == Some("PONYTAIL:OFF") + }); + if turned_off { + result.outputs.push(HookOutput { + extension_id: extension.id.clone(), + event: "SessionStart".into(), + additional_context: Some(String::new()), + system_message: None, + }); + continue; + } + let refresh = HookEvent::SessionStart { + session_id: event.session_id(), + reason: "resume", + }; + for binding in extension + .hooks + .iter() + .filter(|binding| binding.event == refresh.name()) + .filter(|binding| matcher_applies(binding.matcher.as_deref(), &refresh)) + { + result.invoked.insert(extension.id.clone()); + match run_hook( + extension, + binding, + &refresh, + ( + &package_root, + &self.root.join("data").join(&extension.id), + project_root, + ), + model, + cancel, + ) { + Ok(mut output) => { + output.additional_context.get_or_insert_with(String::new); + result.outputs.push(output); + } + Err(error) => { + result.errors.entry(extension.id.clone()).or_insert(error); + break; + } + } + } + } + } + } + result + } + + pub(crate) fn status_for(&self, events: &[HookEvent]) -> Option { + let statuses = self + .extensions + .iter() + .filter(|extension| extension.enabled) + .flat_map(|extension| { + extension.hooks.iter().filter_map(|binding| { + events + .iter() + .any(|event| { + binding.event == event.name() + && matcher_applies(binding.matcher.as_deref(), event) + }) + .then(|| binding.status_message.clone()) + .flatten() + }) + }) + .collect::>(); + (!statuses.is_empty()).then(|| statuses.into_iter().collect::>().join(" · ")) + } + + pub(crate) fn has_hooks_for(&self, events: &[HookEvent]) -> bool { + self.extensions + .iter() + .filter(|extension| extension.enabled) + .any(|extension| { + extension.hooks.iter().any(|binding| { + events.iter().any(|event| { + binding.event == event.name() + && matcher_applies(binding.matcher.as_deref(), event) + }) + }) + }) + } + + fn validate_skill_conflicts(&self) -> Result<(), String> { + let mut names = BTreeMap::::new(); + if let Some(home) = std::env::var_os("HOME") { + for skill in + crate::agent::discover_agent_skills(&PathBuf::from(home).join(".agents/skills")) + { + names.insert(skill.name, "standard agent skills".into()); + } + } + for (id, root) in self.enabled_skill_roots()? { + for skill in crate::agent::discover_agent_skills(&root) { + if let Some(owner) = names.insert(skill.name.clone(), id.clone()) { + return Err(format!( + "Skill {} is declared by both {owner} and {id}", + skill.name + )); + } + } + } + Ok(()) + } + + fn package_root(&self, id: &str) -> PathBuf { + self.root.join("packages").join(id).join("current") + } + + fn save(&self) -> Result<(), String> { + fs::create_dir_all(&self.root).map_err(|error| error.to_string())?; + let path = self.root.join("registry.json"); + let temporary = self.root.join("registry.json.tmp"); + let bytes = serde_json::to_vec_pretty(self).map_err(|error| error.to_string())?; + fs::write(&temporary, bytes) + .map_err(|error| format!("Could not write {}: {error}", temporary.display()))?; + fs::rename(&temporary, &path) + .map_err(|error| format!("Could not publish {}: {error}", path.display())) + } +} + +#[derive(Clone, Debug)] +pub(crate) enum HookEvent { + SessionStart { + session_id: i32, + reason: &'static str, + }, + UserPromptSubmit { + session_id: i32, + prompt: String, + }, + SubagentStart { + session_id: i32, + round: usize, + }, +} + +impl HookEvent { + pub(crate) fn name(&self) -> &'static str { + match self { + Self::SessionStart { .. } => "SessionStart", + Self::UserPromptSubmit { .. } => "UserPromptSubmit", + Self::SubagentStart { .. } => "SubagentStart", + } + } + + fn session_id(&self) -> i32 { + match self { + Self::SessionStart { session_id, .. } + | Self::UserPromptSubmit { session_id, .. } + | Self::SubagentStart { session_id, .. } => *session_id, + } + } + + fn matcher_value(&self) -> &str { + match self { + Self::SessionStart { reason, .. } => reason, + Self::SubagentStart { .. } => "ralph", + Self::UserPromptSubmit { .. } => "", + } + } + + fn payload(&self, project_root: &Path, model: &str) -> Value { + let mut payload = serde_json::json!({ + "hook_event_name": self.name(), + "event_name": self.name(), + "session_id": self.session_id().to_string(), + "cwd": project_root, + "project_dir": project_root, + "model": model, + "timestamp": SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_secs(), + }); + match self { + Self::SessionStart { reason, .. } => { + payload["source"] = Value::String((*reason).into()) + } + Self::UserPromptSubmit { prompt, .. } => { + payload["prompt"] = Value::String(prompt.clone()) + } + Self::SubagentStart { round, .. } => { + payload["agent_type"] = Value::String("ralph".into()); + payload["round"] = Value::from(*round as u64); + } + } + payload + } +} + +#[derive(Clone, Debug)] +pub(crate) struct HookOutput { + pub(crate) extension_id: String, + pub(crate) event: String, + pub(crate) additional_context: Option, + pub(crate) system_message: Option, +} + +#[derive(Clone, Debug, Default)] +pub(crate) struct HookBatchResult { + pub(crate) outputs: Vec, + pub(crate) errors: BTreeMap, + invoked: BTreeSet, +} + +impl HookBatchResult { + pub(crate) fn status(&self) -> Option { + let statuses = self + .outputs + .iter() + .filter_map(|output| output.system_message.as_deref()) + .collect::>(); + (!statuses.is_empty()).then(|| statuses.join(" · ")) + } +} + +pub(crate) fn wrap_context(id: &str, event: &str, context: &str) -> String { + format!("{CONTEXT_PREFIX}{id} {event}]\n{context}") +} + +pub(crate) fn context_identity(content: &str) -> Option<(&str, &str)> { + let header = content.strip_prefix(CONTEXT_PREFIX)?.split_once("]\n")?.0; + let (id, event) = header.split_once(' ')?; + Some((id, event)) +} + +#[derive(Deserialize)] +struct RawManifest { + name: String, + version: String, + #[serde(default)] + description: String, + #[serde(default)] + author: Option, + #[serde(default)] + skills: Option, + #[serde(default)] + hooks: Option, + #[serde(default, alias = "schemaVersion")] + schema_version: Option, + #[serde(flatten)] + extra: BTreeMap, +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum RawAuthor { + Name(String), + Object { name: String }, +} + +impl RawAuthor { + fn name(self) -> String { + match self { + Self::Name(name) | Self::Object { name } => name, + } + } +} + +#[derive(Deserialize)] +struct RawHookManifest { + hooks: BTreeMap>, +} + +#[derive(Deserialize)] +struct RawHookGroup { + #[serde(default)] + matcher: Option, + hooks: Vec, +} + +#[derive(Deserialize)] +struct RawHook { + #[serde(rename = "type")] + kind: String, + command: String, + #[serde(default = "default_timeout")] + timeout: u64, + #[serde(default, rename = "statusMessage")] + status_message: Option, +} + +fn default_timeout() -> u64 { + 5 +} + +fn validate_package( + root: &Path, + source_url: &str, + requested_ref: Option, + commit: git2::Oid, +) -> Result { + let root = root.canonicalize().map_err(|error| error.to_string())?; + let manifest_path = resolve_inside(&root, &root.join(MANIFEST_PATH))?; + let raw = fs::read(&manifest_path) + .map_err(|error| format!("Could not read {}: {error}", manifest_path.display()))?; + let manifest: RawManifest = serde_json::from_slice(&raw) + .map_err(|error| format!("Invalid extension manifest: {error}"))?; + if manifest.schema_version.is_some_and(|version| version != 1) { + return Err(format!( + "Extension manifest schema {} is unsupported; expected 1", + manifest.schema_version.unwrap() + )); + } + if !valid_id(&manifest.name) { + return Err( + "Extension name must be a lowercase ID using letters, digits, and hyphens".into(), + ); + } + if manifest.version.trim().is_empty() || manifest.version.len() > 64 { + return Err("Extension version must be 1 to 64 characters".into()); + } + for capability in [ + "commands", + "agents", + "mcpServers", + "mcp_servers", + "apps", + "scripts", + ] { + if manifest.extra.contains_key(capability) { + return Err(format!( + "Extension capability {capability} is executable and unsupported by DS4Server" + )); + } + } + let skills_path = manifest + .skills + .as_deref() + .map(|path| validate_relative(path, "skills")) + .transpose()?; + let mut skill_count = 0; + if let Some(path) = &skills_path { + let skills = resolve_inside(&root, &root.join(path))?; + reject_symlink_escapes(&root, &skills)?; + let discovered = crate::agent::discover_agent_skills(&skills); + skill_count = discovered.len(); + if skill_count == 0 { + return Err(format!( + "Extension skill directory {} has no valid skills", + skills.display() + )); + } + } + let hooks_path = manifest + .hooks + .as_deref() + .map(|path| validate_relative(path, "hooks")) + .transpose()?; + if skills_path.is_none() && hooks_path.is_none() { + return Err("Extension declares neither skills nor lifecycle hooks".into()); + } + let hooks = hooks_path + .as_ref() + .map(|path| parse_hooks(&root, path)) + .transpose()? + .unwrap_or_default(); + Ok(InstalledExtension { + id: manifest.name.clone(), + name: manifest.name, + version: manifest.version, + description: manifest.description, + author: manifest.author.map(RawAuthor::name).unwrap_or_default(), + enabled: false, + trusted: false, + source_url: source_url.into(), + requested_ref, + resolved_commit: commit.to_string(), + skills_path, + skill_count, + hooks_path, + hooks, + last_error: None, + }) +} + +fn parse_hooks(root: &Path, relative: &str) -> Result, String> { + let path = resolve_inside(root, &root.join(relative))?; + let raw = + fs::read(&path).map_err(|error| format!("Could not read {}: {error}", path.display()))?; + let manifest: RawHookManifest = + serde_json::from_slice(&raw).map_err(|error| format!("Invalid hook manifest: {error}"))?; + let mut bindings = Vec::new(); + for (event, groups) in manifest.hooks { + if !matches!( + event.as_str(), + "SessionStart" | "UserPromptSubmit" | "SubagentStart" + ) { + return Err(format!("Hook event {event} is unsupported by DS4Server")); + } + for group in groups { + if let Some(matcher) = group.matcher.as_deref() { + Regex::new(matcher).map_err(|error| format!("Invalid {event} matcher: {error}"))?; + } + for hook in group.hooks { + if hook.kind != "command" { + return Err(format!("Hook type {} is unsupported", hook.kind)); + } + if hook.timeout == 0 || hook.timeout > MAX_HOOK_TIMEOUT { + return Err(format!( + "Hook timeout must be between 1 and {MAX_HOOK_TIMEOUT} seconds" + )); + } + let argv = shlex::split(&hook.command) + .ok_or_else(|| format!("Hook command has invalid quoting: {}", hook.command))?; + if argv.is_empty() { + return Err("Hook command is empty".into()); + } + validate_command_paths(root, &argv)?; + bindings.push(HookBinding { + event: event.clone(), + matcher: group.matcher.clone(), + argv, + timeout_seconds: hook.timeout, + status_message: hook.status_message, + }); + } + } + } + Ok(bindings) +} + +fn run_hook( + extension: &InstalledExtension, + binding: &HookBinding, + event: &HookEvent, + roots: (&Path, &Path, &Path), + model: &str, + cancel: &AtomicBool, +) -> Result { + if cancel.load(Ordering::Relaxed) { + return Err(format!("{} hook was cancelled", event.name())); + } + let (package_root, data_root, project_root) = roots; + let package_root = package_root + .canonicalize() + .map_err(|error| error.to_string())?; + let project_root = project_root + .canonicalize() + .map_err(|error| format!("Could not resolve project root: {error}"))?; + let session_data = data_root + .join("sessions") + .join(event.session_id().to_string()); + let config = data_root.join("config"); + let home = data_root.join("home"); + let temporary = data_root.join("tmp"); + for directory in [&session_data, &config, &home, &temporary] { + fs::create_dir_all(directory).map_err(|error| error.to_string())?; + } + let argv = expand_argv(&binding.argv, &package_root, &session_data, data_root)?; + let mut event_payload = event.payload(&project_root, model); + if extension.id == "ponytail" + && let HookEvent::UserPromptSubmit { prompt, .. } = event + && prompt.trim() == "/ponytail status" + { + // Ponytail 4.9 implements status as the bare command. DS4Server keeps + // the user's original prompt intact and adapts only the hook payload. + event_payload["prompt"] = Value::String("/ponytail".into()); + } + let payload = serde_json::to_vec(&event_payload).map_err(|error| error.to_string())?; + if payload.len() > MAX_HOOK_INPUT { + return Err("Hook event exceeds the input limit".into()); + } + let mut command = Command::new(&argv[0]); + command + .args(&argv[1..]) + .current_dir(&project_root) + .env_clear() + .env("PATH", std::env::var_os("PATH").unwrap_or_default()) + .env("HOME", &home) + .env("TMPDIR", &temporary) + .env("XDG_CONFIG_HOME", &config) + .env("CLAUDE_CONFIG_DIR", data_root.join("claude")) + .env("CLAUDE_PLUGIN_ROOT", &package_root) + .env("PLUGIN_DATA", &session_data) + .env("DS4SERVER_PLUGIN_ROOT", &package_root) + .env("DS4SERVER_PLUGIN_DATA", data_root) + .env("DS4SERVER_PROJECT_ROOT", &project_root) + .env("DS4SERVER_SESSION_ID", event.session_id().to_string()) + .env("DS4SERVER_EVENT_NAME", event.name()) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .process_group(0); + if extension.id == "ponytail" + && matches!(event, HookEvent::SessionStart { .. }) + && let Ok(mode) = fs::read_to_string(session_data.join(".ponytail-active")) + && matches!(mode.trim(), "off" | "lite" | "full" | "ultra") + { + command.env("PONYTAIL_DEFAULT_MODE", mode.trim()); + } + let mut child = command.spawn().map_err(|error| { + format!( + "Could not start {} hook executable {}: {error}. Install it or disable the extension.", + event.name(), + argv[0] + ) + })?; + child + .stdin + .take() + .ok_or_else(|| "Hook stdin is unavailable".to_owned())? + .write_all(&payload) + .map_err(|error| format!("Could not write hook input: {error}"))?; + let stdout = child + .stdout + .take() + .ok_or_else(|| "Hook stdout is unavailable".to_owned())?; + let stderr = child + .stderr + .take() + .ok_or_else(|| "Hook stderr is unavailable".to_owned())?; + let stdout = thread::spawn(move || read_bounded(stdout)); + let stderr = thread::spawn(move || read_bounded(stderr)); + let deadline = Instant::now() + Duration::from_secs(binding.timeout_seconds); + let status = loop { + if cancel.load(Ordering::Relaxed) { + terminate_process_group(&mut child); + return Err(format!("{} hook was cancelled", event.name())); + } + if Instant::now() >= deadline { + terminate_process_group(&mut child); + return Err(format!( + "{} hook exceeded its {} second timeout", + event.name(), + binding.timeout_seconds + )); + } + match child.try_wait() { + Ok(Some(status)) => break status, + Ok(None) => thread::sleep(Duration::from_millis(20)), + Err(error) => { + terminate_process_group(&mut child); + return Err(format!("Could not wait for {} hook: {error}", event.name())); + } + } + }; + let stdout = stdout + .join() + .map_err(|_| "Hook stdout reader panicked".to_owned())??; + let stderr = stderr + .join() + .map_err(|_| "Hook stderr reader panicked".to_owned())??; + if !status.success() { + return Err(format!( + "{} hook exited with {}{}", + event.name(), + status, + if stderr.trim().is_empty() { + String::new() + } else { + format!(": {}", stderr.trim()) + } + )); + } + parse_hook_output(&extension.id, event.name(), stdout) +} + +fn parse_hook_output(id: &str, event: &str, stdout: String) -> Result { + let stdout = stdout.trim(); + if stdout.is_empty() { + return Ok(HookOutput { + extension_id: id.into(), + event: event.into(), + additional_context: None, + system_message: None, + }); + } + let Ok(value) = serde_json::from_str::(stdout) else { + if stdout.starts_with('{') || stdout.starts_with('[') { + return Err(format!("{event} hook returned malformed JSON")); + } + return Ok(HookOutput { + extension_id: id.into(), + event: event.into(), + additional_context: Some(stdout.into()), + system_message: None, + }); + }; + let object = value + .as_object() + .ok_or_else(|| format!("{event} hook JSON output must be an object"))?; + let specific = object.get("hookSpecificOutput").and_then(Value::as_object); + if let Some(name) = specific + .and_then(|specific| specific.get("hookEventName")) + .and_then(Value::as_str) + && name != event + { + return Err(format!("{event} hook returned output for {name}")); + } + let additional_context = specific + .and_then(|specific| specific.get("additionalContext")) + .or_else(|| object.get("additionalContext")) + .map(|value| { + value + .as_str() + .map(str::to_owned) + .ok_or_else(|| format!("{event} additionalContext must be a string")) + }) + .transpose()?; + let system_message = object + .get("systemMessage") + .map(|value| { + value + .as_str() + .map(str::to_owned) + .ok_or_else(|| format!("{event} systemMessage must be a string")) + }) + .transpose()?; + Ok(HookOutput { + extension_id: id.into(), + event: event.into(), + additional_context, + system_message, + }) +} + +fn read_bounded(mut reader: impl Read) -> Result { + let mut bytes = Vec::new(); + reader + .by_ref() + .take((MAX_HOOK_OUTPUT + 1) as u64) + .read_to_end(&mut bytes) + .map_err(|error| error.to_string())?; + if bytes.len() > MAX_HOOK_OUTPUT { + return Err(format!("Hook output exceeds {MAX_HOOK_OUTPUT} bytes")); + } + String::from_utf8(bytes).map_err(|_| "Hook output is not valid UTF-8".into()) +} + +fn terminate_process_group(child: &mut std::process::Child) { + let Ok(group) = i32::try_from(child.id()).map(|pid| -pid) else { + let _ = child.kill(); + let _ = child.wait(); + return; + }; + // SAFETY: `group` is the freshly spawned child's process group and the + // signal constants are valid on the supported Unix platform. + unsafe { + libc::kill(group, libc::SIGTERM); + } + let deadline = Instant::now() + Duration::from_millis(250); + while Instant::now() < deadline { + if child.try_wait().ok().flatten().is_some() { + return; + } + thread::sleep(Duration::from_millis(10)); + } + // SAFETY: same process group and platform invariants as above. + unsafe { + libc::kill(group, libc::SIGKILL); + } + let _ = child.wait(); +} + +fn matcher_applies(matcher: Option<&str>, event: &HookEvent) -> bool { + matcher.is_none_or(|matcher| { + Regex::new(matcher).is_ok_and(|matcher| matcher.is_match(event.matcher_value())) + }) +} + +fn expand_argv( + argv: &[String], + package_root: &Path, + session_data: &Path, + data_root: &Path, +) -> Result, String> { + argv.iter() + .map(|argument| { + let expanded = argument + .replace("${CLAUDE_PLUGIN_ROOT}", &package_root.to_string_lossy()) + .replace("${DS4SERVER_PLUGIN_ROOT}", &package_root.to_string_lossy()) + .replace("${PLUGIN_DATA}", &session_data.to_string_lossy()) + .replace("${DS4SERVER_PLUGIN_DATA}", &data_root.to_string_lossy()); + if expanded.contains('$') { + return Err(format!( + "Hook argument contains an unsupported variable: {argument}" + )); + } + Ok(expanded) + }) + .collect() +} + +fn validate_command_paths(root: &Path, argv: &[String]) -> Result<(), String> { + for argument in argv { + let without_variables = [ + "${CLAUDE_PLUGIN_ROOT}", + "${DS4SERVER_PLUGIN_ROOT}", + "${PLUGIN_DATA}", + "${DS4SERVER_PLUGIN_DATA}", + ] + .iter() + .fold(argument.clone(), |value, variable| { + value.replace(variable, "") + }); + if without_variables.contains('$') { + return Err(format!( + "Hook command uses an unsupported variable: {argument}" + )); + } + if argument.contains("${CLAUDE_PLUGIN_ROOT}") + || argument.contains("${DS4SERVER_PLUGIN_ROOT}") + { + let expanded = argument + .replace("${CLAUDE_PLUGIN_ROOT}", &root.to_string_lossy()) + .replace("${DS4SERVER_PLUGIN_ROOT}", &root.to_string_lossy()); + resolve_inside(root, Path::new(&expanded))?; + } + } + Ok(()) +} + +fn ponytail_mode_command(event: &HookEvent) -> bool { + let HookEvent::UserPromptSubmit { prompt, .. } = event else { + return false; + }; + let mut words = prompt.split_whitespace(); + matches!(words.next(), Some("/ponytail") | Some("$ponytail")) + && !matches!(words.next(), Some("default")) +} + +fn validate_relative(path: &str, label: &str) -> Result { + let path = Path::new(path); + if path.is_absolute() + || path.components().any(|component| { + matches!( + component, + Component::ParentDir | Component::RootDir | Component::Prefix(_) + ) + }) + { + return Err(format!( + "Extension {label} path must stay inside the package" + )); + } + let normalized = path + .components() + .filter_map(|component| match component { + Component::Normal(value) => Some(value.to_string_lossy()), + Component::CurDir => None, + _ => None, + }) + .collect::>() + .join("/"); + if normalized.is_empty() { + return Err(format!("Extension {label} path is empty")); + } + Ok(normalized) +} + +fn resolve_inside(root: &Path, path: &Path) -> Result { + let root = root.canonicalize().map_err(|error| error.to_string())?; + let path = path + .canonicalize() + .map_err(|error| format!("Could not resolve {}: {error}", path.display()))?; + if !path.starts_with(&root) { + return Err(format!( + "Path escapes the extension package: {}", + path.display() + )); + } + Ok(path) +} + +fn reject_symlink_escapes(package_root: &Path, path: &Path) -> Result<(), String> { + for entry in fs::read_dir(path).map_err(|error| error.to_string())? { + let entry = entry.map_err(|error| error.to_string())?; + let metadata = fs::symlink_metadata(entry.path()).map_err(|error| error.to_string())?; + if metadata.file_type().is_symlink() { + resolve_inside(package_root, &entry.path())?; + } else if metadata.is_dir() { + reject_symlink_escapes(package_root, &entry.path())?; + } + } + Ok(()) +} + +fn valid_id(id: &str) -> bool { + (1..=64).contains(&id.len()) + && !id.starts_with('-') + && !id.ends_with('-') + && !id.contains("--") + && id + .bytes() + .all(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit() || byte == b'-') +} + +fn validate_source_url(source: &str) -> Result<(), String> { + let url = url::Url::parse(source).map_err(|error| format!("Invalid extension URL: {error}"))?; + if url.scheme() != "https" + || url.host_str().is_none() + || !url.username().is_empty() + || url.password().is_some() + { + return Err( + "Extension source must be an HTTPS Git URL without embedded credentials".into(), + ); + } + Ok(()) +} + +fn checkout_requested( + repository: &git2::Repository, + requested: Option<&str>, +) -> Result { + let object = if let Some(requested) = requested { + if requested.contains(['\0', '\n', '\r']) || requested.starts_with('-') { + return Err("Extension ref contains invalid characters".into()); + } + [ + requested.to_owned(), + format!("refs/tags/{requested}"), + format!("refs/remotes/origin/{requested}"), + ] + .into_iter() + .find_map(|candidate| repository.revparse_single(&candidate).ok()) + .ok_or_else(|| format!("Extension ref {requested} was not found"))? + } else { + repository + .head() + .and_then(|head| head.peel(git2::ObjectType::Commit)) + .map_err(|error| format!("Could not resolve extension HEAD: {error}"))? + }; + let commit = object + .peel_to_commit() + .map_err(|error| format!("Extension ref is not a commit: {error}"))?; + repository + .checkout_tree( + commit.as_object(), + Some(git2::build::CheckoutBuilder::new().force()), + ) + .map_err(|error| format!("Could not check out extension ref: {error}"))?; + repository + .set_head_detached(commit.id()) + .map_err(|error| format!("Could not pin extension commit: {error}"))?; + Ok(commit.id()) +} + +fn promote_package( + root: &Path, + id: &str, + staging: &Path, + backup_label: Option<&str>, +) -> Result, String> { + let directory = root.join("packages").join(id); + fs::create_dir_all(&directory).map_err(|error| error.to_string())?; + let current = directory.join("current"); + let backup = backup_label.map(|label| temporary_path(&directory, label)); + if current.exists() { + let backup = backup + .as_ref() + .ok_or_else(|| "Extension package already exists".to_owned())?; + fs::rename(¤t, backup).map_err(|error| error.to_string())?; + } + if let Err(error) = fs::rename(staging, ¤t) { + if let Some(backup) = &backup { + let _ = fs::rename(backup, ¤t); + } + return Err(format!("Could not activate extension package: {error}")); + } + Ok(backup) +} + +fn rollback_package(root: &Path, id: &str, backup: Option<&Path>) { + let current = root.join("packages").join(id).join("current"); + if current.exists() { + let _ = fs::remove_dir_all(¤t); + } + if let Some(backup) = backup { + let _ = fs::rename(backup, current); + } +} + +fn temporary_path(parent: &Path, label: &str) -> PathBuf { + let stamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + parent.join(format!(".{label}-{}-{stamp}", std::process::id())) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::os::unix::fs::{PermissionsExt, symlink}; + use std::sync::Arc; + + fn fixture(name: &str) -> PathBuf { + let root = std::env::temp_dir().join(format!( + "ds4-extension-{name}-{}-{}", + std::process::id(), + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos() + )); + fs::create_dir_all(root.join(".codex-plugin")).unwrap(); + fs::create_dir_all(root.join("skills/demo")).unwrap(); + fs::create_dir_all(root.join("hooks")).unwrap(); + fs::write( + root.join("skills/demo/SKILL.md"), + "---\nname: demo\ndescription: Demonstrate extension skills\n---\nInstructions\n", + ) + .unwrap(); + fs::write(root.join("hooks/hook"), "fixture hook").unwrap(); + fs::write( + root.join("hooks/hooks.json"), + r#"{"hooks":{"SessionStart":[{"matcher":"startup|resume|compact","hooks":[{"type":"command","command":"${CLAUDE_PLUGIN_ROOT}/hooks/hook","timeout":5}]}]}}"#, + ) + .unwrap(); + fs::write( + root.join(MANIFEST_PATH), + r#"{"name":"fixture","version":"1.0.0","description":"Fixture","author":{"name":"DS4"},"skills":"./skills","hooks":"./hooks/hooks.json"}"#, + ) + .unwrap(); + root + } + + #[test] + fn manifest_and_paths_are_validated_without_following_escapes() { + let root = fixture("manifest"); + let package = validate_package( + &root, + "https://example.com/fixture.git", + None, + git2::Oid::ZERO_SHA1, + ) + .unwrap(); + assert_eq!(package.id, "fixture"); + assert_eq!(package.skill_count, 1); + assert_eq!(package.hook_names(), "SessionStart"); + + fs::write( + root.join(MANIFEST_PATH), + r#"{"name":"fixture","version":"1","schemaVersion":2,"skills":"../outside"}"#, + ) + .unwrap(); + assert!( + validate_package(&root, "https://example.com/x", None, git2::Oid::ZERO_SHA1) + .unwrap_err() + .contains("schema 2") + ); + fs::write( + root.join(MANIFEST_PATH), + r#"{"name":"fixture","version":"1","skills":"/tmp"}"#, + ) + .unwrap(); + assert!( + validate_package(&root, "https://example.com/x", None, git2::Oid::ZERO_SHA1) + .unwrap_err() + .contains("inside") + ); + fs::remove_dir_all(root).unwrap(); + } + + #[test] + fn hook_output_is_bounded_typed_and_event_matched() { + let output = parse_hook_output( + "fixture", + "SessionStart", + r#"{"systemMessage":"ready","hookSpecificOutput":{"hookEventName":"SessionStart","additionalContext":"rules"}}"#.into(), + ) + .unwrap(); + assert_eq!(output.additional_context.as_deref(), Some("rules")); + assert_eq!(output.system_message.as_deref(), Some("ready")); + assert!( + parse_hook_output( + "fixture", + "SessionStart", + r#"{"hookSpecificOutput":{"hookEventName":"UserPromptSubmit"}}"#.into(), + ) + .is_err() + ); + assert!(read_bounded(&b"x"[..]).is_ok()); + assert!(read_bounded(&vec![b'x'; MAX_HOOK_OUTPUT + 1][..]).is_err()); + } + + #[test] + fn direct_hook_runner_isolated_environment_and_handles_failures() { + let root = fixture("runner"); + let script = root.join("hooks/run.sh"); + fs::write( + &script, + "#!/bin/sh\nread input\nprintf '{\"hookSpecificOutput\":{\"hookEventName\":\"SessionStart\",\"additionalContext\":\"%s|%s|%s\"}}' \"$PLUGIN_DATA\" \"$HOME\" \"$DS4SERVER_EVENT_NAME\"\n", + ) + .unwrap(); + let mut permissions = fs::metadata(&script).unwrap().permissions(); + permissions.set_mode(0o755); + fs::set_permissions(&script, permissions).unwrap(); + let package = InstalledExtension { + id: "fixture".into(), + name: "fixture".into(), + version: "1".into(), + description: String::new(), + author: String::new(), + enabled: true, + trusted: true, + source_url: "https://example.com/x".into(), + requested_ref: None, + resolved_commit: "0".into(), + skills_path: None, + skill_count: 0, + hooks_path: None, + hooks: Vec::new(), + last_error: None, + }; + let binding = HookBinding { + event: "SessionStart".into(), + matcher: None, + argv: vec!["${CLAUDE_PLUGIN_ROOT}/hooks/run.sh".into()], + timeout_seconds: 1, + status_message: None, + }; + let data = root.with_extension("data"); + let output = run_hook( + &package, + &binding, + &HookEvent::SessionStart { + session_id: 7, + reason: "startup", + }, + (&root, &data, &root), + "fixture-model", + &AtomicBool::new(false), + ) + .unwrap(); + let context = output.additional_context.unwrap(); + assert!(context.contains("sessions/7")); + assert!(context.contains("/home|SessionStart")); + fs::remove_dir_all(root).unwrap(); + fs::remove_dir_all(data).unwrap(); + } + + fn executable(path: &Path, body: &str) { + fs::write(path, body).unwrap(); + let mut permissions = fs::metadata(path).unwrap().permissions(); + permissions.set_mode(0o755); + fs::set_permissions(path, permissions).unwrap(); + } + + fn test_extension(argv: Vec, timeout_seconds: u64) -> InstalledExtension { + InstalledExtension { + id: "fixture".into(), + name: "fixture".into(), + version: "1".into(), + description: String::new(), + author: String::new(), + enabled: true, + trusted: true, + source_url: "https://example.com/x".into(), + requested_ref: None, + resolved_commit: "0".into(), + skills_path: None, + skill_count: 0, + hooks_path: None, + hooks: vec![HookBinding { + event: "SessionStart".into(), + matcher: None, + argv, + timeout_seconds, + status_message: Some("fixture status".into()), + }], + last_error: None, + } + } + + #[test] + fn registry_persists_enable_disable_errors_and_uninstall() { + let root = fixture("registry").with_extension("registry-data"); + let package = root.join("packages/fixture/current"); + fs::create_dir_all(&package).unwrap(); + fs::create_dir_all(root.join("data/fixture")).unwrap(); + let mut registry = ExtensionRegistry::empty(&root); + registry.extensions.push(test_extension(Vec::new(), 1)); + registry.extensions[0].enabled = false; + registry.save().unwrap(); + assert!(!ExtensionRegistry::load(&root).unwrap().extensions[0].enabled); + assert!( + ExtensionRegistry::set_enabled(&root, "fixture", true) + .unwrap() + .extensions[0] + .enabled + ); + assert!( + !ExtensionRegistry::set_enabled(&root, "fixture", false) + .unwrap() + .extensions[0] + .enabled + ); + + let mut duplicate: Value = + serde_json::from_slice(&fs::read(root.join("registry.json")).unwrap()).unwrap(); + let copy = duplicate["extensions"][0].clone(); + duplicate["extensions"].as_array_mut().unwrap().push(copy); + fs::write( + root.join("registry.json"), + serde_json::to_vec(&duplicate).unwrap(), + ) + .unwrap(); + assert!( + ExtensionRegistry::load(&root) + .unwrap_err() + .contains("duplicate") + ); + registry.save().unwrap(); + + let removed = ExtensionRegistry::uninstall(&root, "fixture").unwrap(); + assert!(removed.extensions.is_empty()); + assert!(!root.join("packages/fixture").exists()); + assert!(!root.join("data/fixture").exists()); + fs::remove_dir_all(root).unwrap(); + } + + #[test] + fn duplicate_skill_names_and_update_rollback_are_deterministic() { + let root = fixture("conflicts").with_extension("registry"); + let mut registry = ExtensionRegistry::empty(&root); + for id in ["first", "second"] { + let skills = root.join(format!("packages/{id}/current/skills/shared")); + fs::create_dir_all(&skills).unwrap(); + fs::write( + skills.join("SKILL.md"), + "---\nname: shared\ndescription: Shared fixture skill\n---\nRules\n", + ) + .unwrap(); + let mut extension = test_extension(Vec::new(), 1); + extension.id = id.into(); + extension.name = id.into(); + extension.skills_path = Some("skills".into()); + extension.skill_count = 1; + registry.extensions.push(extension); + } + assert!( + registry + .validate_skill_conflicts() + .unwrap_err() + .contains("both first and second") + ); + + let current = root.join("packages/update/current"); + fs::create_dir_all(¤t).unwrap(); + fs::write(current.join("version"), "old").unwrap(); + let staging = root.join("staging"); + fs::create_dir_all(&staging).unwrap(); + fs::write(staging.join("version"), "new").unwrap(); + let backup = promote_package(&root, "update", &staging, Some("previous")) + .unwrap() + .unwrap(); + assert_eq!(fs::read_to_string(current.join("version")).unwrap(), "new"); + rollback_package(&root, "update", Some(&backup)); + assert_eq!(fs::read_to_string(current.join("version")).unwrap(), "old"); + fs::remove_dir_all(root).unwrap(); + } + + #[test] + fn symlink_and_variable_escapes_are_rejected() { + let root = fixture("escapes"); + let outside = root.with_extension("outside"); + fs::create_dir_all(&outside).unwrap(); + symlink(&outside, root.join("skills/escape")).unwrap(); + assert!(reject_symlink_escapes(&root, &root.join("skills")).is_err()); + assert!( + validate_command_paths( + &root, + &["${CLAUDE_PLUGIN_ROOT}/hooks/hook$UNTRUSTED".into()] + ) + .unwrap_err() + .contains("unsupported variable") + ); + fs::remove_dir_all(root).unwrap(); + fs::remove_dir_all(outside).unwrap(); + } + + #[test] + fn hook_failures_are_bounded_actionable_and_non_shell() { + let root = fixture("failures"); + let data = root.with_extension("data"); + let event = HookEvent::SessionStart { + session_id: 9, + reason: "startup", + }; + let run = |script: &str, timeout, cancel: &AtomicBool| { + let path = root.join("hooks/failure.sh"); + executable(&path, script); + let extension = test_extension( + vec!["${CLAUDE_PLUGIN_ROOT}/hooks/failure.sh".into()], + timeout, + ); + run_hook( + &extension, + &extension.hooks[0], + &event, + (&root, &data, &root), + "model", + cancel, + ) + }; + assert!( + run( + "#!/bin/sh\necho bad >&2\nexit 7\n", + 1, + &AtomicBool::new(false) + ) + .unwrap_err() + .contains("bad") + ); + assert!( + run("#!/bin/sh\nprintf '\\377'\n", 1, &AtomicBool::new(false)) + .unwrap_err() + .contains("UTF-8") + ); + assert!( + parse_hook_output("fixture", "SessionStart", "{broken".into()) + .unwrap_err() + .contains("malformed JSON") + ); + let missing = test_extension(vec!["ds4server-definitely-missing".into()], 1); + assert!( + run_hook( + &missing, + &missing.hooks[0], + &event, + (&root, &data, &root), + "model", + &AtomicBool::new(false), + ) + .unwrap_err() + .contains("Install it or disable the extension") + ); + assert!( + run("#!/bin/sh\nsleep 3\n", 1, &AtomicBool::new(false)) + .unwrap_err() + .contains("timeout") + ); + + let cancel = Arc::new(AtomicBool::new(false)); + let trigger = Arc::clone(&cancel); + thread::spawn(move || { + thread::sleep(Duration::from_millis(50)); + trigger.store(true, Ordering::Relaxed); + }); + assert!( + run("#!/bin/sh\nsleep 3\n", 2, &cancel) + .unwrap_err() + .contains("cancelled") + ); + + let process_group = root.join("hooks/process-group.sh"); + executable( + &process_group, + "#!/bin/sh\nsleep 30 &\necho $! > \"$PLUGIN_DATA/child.pid\"\nwait\n", + ); + let extension = test_extension( + vec!["${CLAUDE_PLUGIN_ROOT}/hooks/process-group.sh".into()], + 1, + ); + assert!( + run_hook( + &extension, + &extension.hooks[0], + &event, + (&root, &data, &root), + "model", + &AtomicBool::new(false), + ) + .unwrap_err() + .contains("timeout") + ); + let child_pid = fs::read_to_string(data.join("sessions/9/child.pid")) + .unwrap() + .trim() + .parse::() + .unwrap(); + let deadline = Instant::now() + Duration::from_secs(1); + while Instant::now() < deadline { + // SAFETY: signal 0 performs only an existence check for this PID. + if unsafe { libc::kill(child_pid, 0) } != 0 { + break; + } + thread::sleep(Duration::from_millis(10)); + } + // SAFETY: signal 0 performs only an existence check for this PID. + assert_ne!(unsafe { libc::kill(child_pid, 0) }, 0); + fs::remove_dir_all(root).unwrap(); + fs::remove_dir_all(data).unwrap(); + } + + #[test] + fn lifecycle_matchers_preserve_declared_event_order() { + let root = fixture("ordering").with_extension("runtime"); + let package = root.join("packages/fixture/current"); + fs::create_dir_all(package.join("hooks")).unwrap(); + executable( + &package.join("hooks/run.sh"), + "#!/bin/sh\ninput=$(cat)\nevent=$(printf '%s' \"$input\" | sed -n 's/.*\"hook_event_name\":\"\\([^\"]*\\)\".*/\\1/p')\nprintf '{\"hookSpecificOutput\":{\"hookEventName\":\"%s\",\"additionalContext\":\"%s\"}}' \"$event\" \"$event\"\n", + ); + let mut extension = test_extension(vec!["${CLAUDE_PLUGIN_ROOT}/hooks/run.sh".into()], 1); + extension.hooks.push(HookBinding { + event: "UserPromptSubmit".into(), + matcher: None, + ..extension.hooks[0].clone() + }); + let second_package = root.join("packages/second/current/hooks"); + fs::create_dir_all(&second_package).unwrap(); + fs::copy(package.join("hooks/run.sh"), second_package.join("run.sh")).unwrap(); + let mut second = extension.clone(); + second.id = "second".into(); + second.name = "second".into(); + let mut registry = ExtensionRegistry::empty(&root); + registry.extensions.push(extension); + registry.extensions.push(second); + let result = registry.dispatch( + &[ + HookEvent::SessionStart { + session_id: 1, + reason: "startup", + }, + HookEvent::UserPromptSubmit { + session_id: 1, + prompt: "hello".into(), + }, + ], + &package, + "model", + &AtomicBool::new(false), + ); + assert!(result.errors.is_empty(), "{:?}", result.errors); + assert_eq!( + result + .outputs + .iter() + .map(|output| (output.extension_id.as_str(), output.event.as_str())) + .collect::>(), + [ + ("fixture", "SessionStart"), + ("fixture", "UserPromptSubmit"), + ("second", "SessionStart"), + ("second", "UserPromptSubmit"), + ] + ); + fs::remove_dir_all(root).unwrap(); + } + + #[test] + #[ignore = "downloads the pinned Ponytail reference integration"] + fn pinned_ponytail_runs_all_lifecycle_modes() { + let root = fixture("ponytail-live").with_extension("installed"); + let project = root.join("project"); + fs::create_dir_all(&project).unwrap(); + let mut registry = ExtensionRegistry::install( + &root, + "https://github.com/DietrichGebert/ponytail.git", + Some("2ed6c52c9d7e5e56942508591085fd45dea277d3"), + ) + .unwrap(); + assert_eq!(registry.extensions[0].version, "4.9.0"); + assert_eq!(registry.extensions[0].skill_count, 6); + registry = ExtensionRegistry::trust_and_enable(&root, "ponytail").unwrap(); + let cancel = AtomicBool::new(false); + let startup = registry.dispatch( + &[HookEvent::SessionStart { + session_id: 42, + reason: "startup", + }], + &project, + "model", + &cancel, + ); + assert!(startup.errors.is_empty(), "{:?}", startup.errors); + assert!(startup.outputs.iter().any(|output| { + output + .additional_context + .as_deref() + .is_some_and(|context| context.contains("Ponytail")) + })); + for mode in ["lite", "full", "ultra"] { + let switched = registry.dispatch( + &[HookEvent::UserPromptSubmit { + session_id: 42, + prompt: format!("/ponytail {mode}"), + }], + &project, + "model", + &cancel, + ); + assert!(switched.errors.is_empty(), "{:?}", switched.errors); + assert!(switched.outputs.iter().any(|output| { + output.event == "SessionStart" + && output + .additional_context + .as_deref() + .is_some_and(|context| context.contains(mode)) + })); + } + let status = registry.dispatch( + &[HookEvent::UserPromptSubmit { + session_id: 42, + prompt: "/ponytail status".into(), + }], + &project, + "model", + &cancel, + ); + assert!(status.errors.is_empty(), "{:?}", status.errors); + assert!( + status + .status() + .is_some_and(|status| status.contains("ULTRA")), + "{:?}", + status.outputs + ); + let off = registry.dispatch( + &[HookEvent::UserPromptSubmit { + session_id: 42, + prompt: "/ponytail off".into(), + }], + &project, + "model", + &cancel, + ); + assert!(off.errors.is_empty(), "{:?}", off.errors); + assert!( + off.outputs.iter().any(|output| { + output.event == "SessionStart" && output.additional_context.as_deref() == Some("") + }), + "{:?}", + off.outputs + ); + let full = registry.dispatch( + &[HookEvent::UserPromptSubmit { + session_id: 42, + prompt: "/ponytail full".into(), + }], + &project, + "model", + &cancel, + ); + assert!(full.errors.is_empty(), "{:?}", full.errors); + let default = registry.dispatch( + &[HookEvent::UserPromptSubmit { + session_id: 42, + prompt: "/ponytail default lite".into(), + }], + &project, + "model", + &cancel, + ); + assert!(default.errors.is_empty(), "{:?}", default.errors); + assert_eq!( + fs::read_to_string(root.join("data/ponytail/sessions/42/.ponytail-active")) + .unwrap() + .trim(), + "full" + ); + let new_session = registry.dispatch( + &[HookEvent::SessionStart { + session_id: 43, + reason: "startup", + }], + &project, + "model", + &cancel, + ); + assert!(new_session.errors.is_empty(), "{:?}", new_session.errors); + assert!(new_session.outputs.iter().any(|output| { + output + .additional_context + .as_deref() + .is_some_and(|context| context.contains("lite")) + })); + for event in [ + HookEvent::SessionStart { + session_id: 42, + reason: "compact", + }, + HookEvent::SessionStart { + session_id: 42, + reason: "resume", + }, + HookEvent::SubagentStart { + session_id: 42, + round: 1, + }, + ] { + let result = registry.dispatch(&[event], &project, "model", &cancel); + assert!(result.errors.is_empty(), "{:?}", result.errors); + assert!(result.outputs.iter().any(|output| { + output + .additional_context + .as_deref() + .is_some_and(|context| context.contains("full")) + })); + } + registry = ExtensionRegistry::set_enabled(&root, "ponytail", false).unwrap(); + assert!( + registry + .dispatch( + &[HookEvent::SessionStart { + session_id: 42, + reason: "resume" + }], + &project, + "model", + &cancel + ) + .outputs + .is_empty() + ); + assert!( + ExtensionRegistry::uninstall(&root, "ponytail") + .unwrap() + .extensions + .is_empty() + ); + fs::remove_dir_all(root).unwrap(); + } +} diff --git a/src/main.rs b/src/main.rs index 437d25a..728ff6b 100644 --- a/src/main.rs +++ b/src/main.rs @@ -10,6 +10,7 @@ mod database; mod dev_brain; mod dsml; mod engine; +mod extensions; mod instructions; mod metrics; mod model;