diff --git a/Cargo.lock b/Cargo.lock index 505c9bc..a6d3969 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -276,6 +276,16 @@ dependencies = [ "windows-link", ] +[[package]] +name = "combine" +version = "4.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba5a308b75df32fe02788e748662718f03fde005016435c444eea572398219fd" +dependencies = [ + "bytes", + "memchr", +] + [[package]] name = "compact_str" version = "0.9.1" @@ -1431,6 +1441,55 @@ version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jni" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5efd9a482cf3a427f00d6b35f14332adc7902ce91efb778580e180ff90fa3498" +dependencies = [ + "cfg-if", + "combine", + "jni-macros", + "jni-sys", + "log", + "simd_cesu8", + "thiserror 2.0.18", + "walkdir", + "windows-link", +] + +[[package]] +name = "jni-macros" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a00109accc170f0bdb141fed3e393c565b6f5e072365c3bd58f5b062591560a3" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "simd_cesu8", + "syn 2.0.118", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn 2.0.118", +] + [[package]] name = "js-sys" version = "0.3.103" @@ -1681,6 +1740,12 @@ dependencies = [ "tempfile", ] +[[package]] +name = "ndk-context" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "27b02d87554356db9e9a873add8782d4ea6e3e58ea071a9adb9a2e8ddb884a8b" + [[package]] name = "new_debug_unreachable" version = "1.0.6" @@ -1775,6 +1840,31 @@ dependencies = [ "libc", ] +[[package]] +name = "objc2" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a12a8ed07aefc768292f076dc3ac8c48f3781c8f2d5851dd3d98950e8c5a89f" +dependencies = [ + "objc2-encode", +] + +[[package]] +name = "objc2-encode" +version = "4.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ef25abbcd74fb2609453eb695bd2f860d389e457f67dc17cafc8b8cbc89d0c33" + +[[package]] +name = "objc2-foundation" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3e0adef53c21f888deb4fa59fc59f7eb17404926ee8a6f59f5df0fd7f9f3272" +dependencies = [ + "bitflags 2.13.0", + "objc2", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -2818,6 +2908,22 @@ version = "0.3.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "703d5c7ef118737c72f1af64ad2f6f8c5e1921f818cdcb97b8fe6fc69bf66214" +[[package]] +name = "simd_cesu8" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94f90157bb87cddf702797c5dadfa0be7d266cdf49e22da2fcaa32eff75b2c33" +dependencies = [ + "rustc_version", + "simdutf8", +] + +[[package]] +name = "simdutf8" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" + [[package]] name = "siphasher" version = "1.0.3" @@ -3644,6 +3750,22 @@ dependencies = [ "string_cache_codegen", ] +[[package]] +name = "webbrowser" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0fc95580916af1e68ff6a7be07446fc5db73ebf71cf092de939bbf5f7e189f72" +dependencies = [ + "core-foundation 0.10.1", + "jni", + "log", + "ndk-context", + "objc2", + "objc2-foundation", + "url", + "web-sys", +] + [[package]] name = "wezterm-bidi" version = "0.2.3" @@ -4184,6 +4306,7 @@ dependencies = [ "tracing-subscriber", "url", "uuid", + "webbrowser", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index ed63a6c..f04e13f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -36,6 +36,7 @@ rmcp = { version = "1.8", default-features = false, features = ["client", "trans chrono = { version = "0.4", features = ["serde"] } tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter"] } +webbrowser = "1" [[bin]] name = "zesdex" diff --git a/src/app/bgbash/control.rs b/src/app/bgbash/control.rs index f2866a5..b60d426 100644 --- a/src/app/bgbash/control.rs +++ b/src/app/bgbash/control.rs @@ -28,11 +28,3 @@ pub fn bash_kill(id: &str) -> anyhow::Result<()> { anyhow::bail!("bash job '{}' not found", id) } } - -pub fn register_bash_job(job: BashJob) -> String { - let id = job.id.clone(); - if let Ok(mut map) = bash_jobs_map().lock() { - map.insert(id.clone(), job); - } - id -} diff --git a/src/app/bgbash/job.rs b/src/app/bgbash/job.rs index 6f19af5..2df9529 100644 --- a/src/app/bgbash/job.rs +++ b/src/app/bgbash/job.rs @@ -5,20 +5,16 @@ use std::io::BufRead; pub struct BashJob { pub id: String, - pub command: String, - pub started_at: i64, pub output_rx: mpsc::Receiver, pub exit_code: Option, - pub handle: Option>, } pub fn spawn_bash_job(command: String) -> BashJob { let id = uuid::Uuid::new_v4().to_string(); - let started_at = chrono::Utc::now().timestamp_millis(); let (output_tx, output_rx) = mpsc::channel::(); let cmd = command.clone(); - let handle = thread::spawn(move || { + let _handle = thread::spawn(move || { let child = Command::new("sh") .arg("-c") .arg(&cmd) @@ -46,11 +42,8 @@ pub fn spawn_bash_job(command: String) -> BashJob { BashJob { id, - command, - started_at, output_rx, exit_code: None, - handle: Some(handle), } } @@ -68,8 +61,4 @@ impl BashJob { Err(_) => None, } } - - pub fn is_running(&self) -> bool { - self.exit_code.is_none() - } } diff --git a/src/app/harness.rs b/src/app/harness.rs index fa111fa..8e40a84 100644 --- a/src/app/harness.rs +++ b/src/app/harness.rs @@ -5,25 +5,9 @@ pub enum Verdict { Escalate, } -impl Verdict { - pub fn is_allowed(&self) -> bool { - matches!(self, Verdict::Allow) - } -} - pub struct Harness; impl Harness { - pub fn classify(_cmd: &str, mode: &super::state::types::AgentMode) -> Verdict { - if mode.auto_approve() { - return Verdict::Allow; - } - if matches!(mode, super::state::types::AgentMode::Plan) { - return Verdict::Block("mutating tools are disabled in Plan mode".to_string()); - } - Verdict::Escalate - } - pub fn gate_tool_call( tool_name: &str, args: &serde_json::Value, @@ -39,6 +23,16 @@ impl Harness { Self::classify(tool_name, mode) } + fn classify(_cmd: &str, mode: &super::state::types::AgentMode) -> Verdict { + if mode.auto_approve() { + return Verdict::Allow; + } + if matches!(mode, super::state::types::AgentMode::Plan) { + return Verdict::Block("mutating tools are disabled in Plan mode".to_string()); + } + Verdict::Escalate + } + fn run_catastrophic_guard( tool_name: &str, args: &serde_json::Value, @@ -73,43 +67,6 @@ impl Harness { } } -pub fn parse_verdict(text: &str) -> Option { - let trimmed = text.trim(); - if let Ok(v) = serde_json::from_str::(trimmed) { - if let Some(verdict) = v.get("verdict").and_then(|v| v.as_str()) { - return match verdict.to_lowercase().as_str() { - "allow" => Some(Verdict::Allow), - "block" => Some(Verdict::Block( - v.get("reason").and_then(|r| r.as_str()).unwrap_or("blocked").to_string() - )), - "escalate" => Some(Verdict::Escalate), - _ => None, - }; - } - } - for line in trimmed.lines() { - let l = line.trim().to_lowercase(); - if l.starts_with("verdict: allow") { - return Some(Verdict::Allow); - } - if l.starts_with("verdict: block") { - let reason = line.split_once(':').map(|x| x.1).unwrap_or("blocked").trim().to_string(); - return Some(Verdict::Block(reason)); - } - } - if trimmed.to_lowercase().contains("allow") { - return Some(Verdict::Allow); - } - if trimmed.to_lowercase().contains("block") { - return Some(Verdict::Block("blocked by classifier".to_string())); - } - None -} - -pub fn classify(_cmd: &str, mode: &super::state::types::AgentMode) -> Verdict { - Harness::classify(_cmd, mode) -} - impl Default for Harness { fn default() -> Self { Harness @@ -122,11 +79,37 @@ mod tests { use crate::app::state::types::AgentMode; use serde_json::json; - #[test] - fn test_verdict_is_allowed() { - assert!(Verdict::Allow.is_allowed()); - assert!(!Verdict::Block("test".to_string()).is_allowed()); - assert!(!Verdict::Escalate.is_allowed()); + fn parse_verdict(text: &str) -> Option { + let trimmed = text.trim(); + if let Ok(v) = serde_json::from_str::(trimmed) { + if let Some(verdict) = v.get("verdict").and_then(|v| v.as_str()) { + return match verdict.to_lowercase().as_str() { + "allow" => Some(Verdict::Allow), + "block" => Some(Verdict::Block( + v.get("reason").and_then(|r| r.as_str()).unwrap_or("blocked").to_string() + )), + "escalate" => Some(Verdict::Escalate), + _ => None, + }; + } + } + for line in trimmed.lines() { + let l = line.trim().to_lowercase(); + if l.starts_with("verdict: allow") { + return Some(Verdict::Allow); + } + if l.starts_with("verdict: block") { + let reason = line.split_once(':').map(|x| x.1).unwrap_or("blocked").trim().to_string(); + return Some(Verdict::Block(reason)); + } + } + if trimmed.to_lowercase().contains("allow") { + return Some(Verdict::Allow); + } + if trimmed.to_lowercase().contains("block") { + return Some(Verdict::Block("blocked by classifier".to_string())); + } + None } #[test] diff --git a/src/app/mcp/manager.rs b/src/app/mcp/manager.rs index a036de0..e0b6d3e 100644 --- a/src/app/mcp/manager.rs +++ b/src/app/mcp/manager.rs @@ -1,5 +1,10 @@ -use serde_json::Value; +use serde_json::{json, Value}; use serde::{Deserialize, Serialize}; +use std::io::{BufRead, BufReader, Write}; + + +const MCP_CONNECT_TIMEOUT_MS: u64 = 20_000; +const MCP_CALL_TIMEOUT_MS: u64 = 60_000; #[derive(Debug, Clone, Serialize, Deserialize)] pub enum McpTransport { @@ -26,43 +31,211 @@ pub struct McpServer { pub tools: Vec, } -impl McpServer { - pub fn new(name: String, transport: McpTransport) -> Self { - McpServer { - name, - transport, - tools: Vec::new(), +#[derive(Debug)] +struct StdioChild { + stdin: std::process::ChildStdin, + stdout: BufReader, + next_id: u64, +} + +impl StdioChild { + fn call(&mut self, method: &str, params: Value) -> anyhow::Result { + self.next_id += 1; + let id = self.next_id; + let req = json!({ + "jsonrpc": "2.0", + "id": id, + "method": method, + "params": params + }); + let mut line = serde_json::to_string(&req)?; + line.push('\n'); + self.stdin.write_all(line.as_bytes())?; + self.stdin.flush()?; + + let mut response_line = String::new(); + let deadline = std::time::Instant::now() + + std::time::Duration::from_millis(MCP_CALL_TIMEOUT_MS); + loop { + if std::time::Instant::now() > deadline { + anyhow::bail!("MCP call timed out after {}ms", MCP_CALL_TIMEOUT_MS); + } + response_line.clear(); + match self.stdout.read_line(&mut response_line) { + Ok(0) => anyhow::bail!("MCP stdio child process closed unexpectedly"), + Ok(_) => { + let trimmed = response_line.trim(); + if trimmed.is_empty() { + continue; + } + let resp: Value = serde_json::from_str(trimmed) + .map_err(|e| anyhow::anyhow!("invalid JSON from MCP server: {}", e))?; + if resp.get("id") == Some(&json!(id)) { + if let Some(err) = resp.get("error") { + anyhow::bail!("MCP error: {}", err); + } + return Ok(resp.get("result").cloned().unwrap_or(Value::Null)); + } + } + Err(e) => anyhow::bail!("MCP stdio read error: {}", e), + } } } } +fn spawn_stdio_child(command: &str, extra_args: &[String]) -> anyhow::Result { + let parts: Vec<&str> = command.split_whitespace().collect(); + let (prog, prog_args) = parts.split_first() + .ok_or_else(|| anyhow::anyhow!("MCP stdio command is empty"))?; + + let mut cmd = std::process::Command::new(prog); + cmd.args(prog_args); + cmd.args(extra_args); + cmd.stdin(std::process::Stdio::piped()); + cmd.stdout(std::process::Stdio::piped()); + cmd.stderr(std::process::Stdio::null()); + + let mut child = cmd.spawn() + .map_err(|e| anyhow::anyhow!("failed to spawn MCP stdio server '{}': {}", command, e))?; + + let stdin = child.stdin.take() + .ok_or_else(|| anyhow::anyhow!("failed to get stdin for MCP server"))?; + let stdout = child.stdout.take() + .ok_or_else(|| anyhow::anyhow!("failed to get stdout for MCP server"))?; + + let mut mcp = StdioChild { + stdin, + stdout: BufReader::new(stdout), + next_id: 0, + }; + + let deadline = std::time::Instant::now() + + std::time::Duration::from_millis(MCP_CONNECT_TIMEOUT_MS); + + let init_result = mcp.call("initialize", json!({ + "protocolVersion": "2024-11-05", + "capabilities": {}, + "clientInfo": { + "name": "zesdex", + "version": "0.1.0" + } + })); + + if std::time::Instant::now() > deadline { + anyhow::bail!("MCP initialize timed out"); + } + + init_result.map_err(|e| anyhow::anyhow!("MCP initialize failed: {}", e))?; + + let _ = mcp.call("notifications/initialized", json!({})); + + Ok(mcp) +} + +fn call_via_stdio(command: &str, extra_args: &[String], tool_name: &str, tool_args: &Value) -> anyhow::Result { + let mut child = spawn_stdio_child(command, extra_args)?; + let result = child.call("tools/call", json!({ + "name": tool_name, + "arguments": tool_args + }))?; + + extract_text_content(&result) +} + +fn call_via_http(url: &str, tool_name: &str, tool_args: &Value) -> anyhow::Result { + let client = reqwest::blocking::Client::builder() + .timeout(std::time::Duration::from_millis(MCP_CALL_TIMEOUT_MS)) + .connect_timeout(std::time::Duration::from_millis(MCP_CONNECT_TIMEOUT_MS)) + .build() + .unwrap_or_else(|_| reqwest::blocking::Client::new()); + + let request_id: u64 = 1; + let body = json!({ + "jsonrpc": "2.0", + "id": request_id, + "method": "tools/call", + "params": { + "name": tool_name, + "arguments": tool_args + } + }); + + let resp = client.post(url) + .header("Content-Type", "application/json") + .json(&body) + .send() + .map_err(|e| anyhow::anyhow!("MCP HTTP request failed: {}", e))?; + + if !resp.status().is_success() { + let status = resp.status(); + let text = resp.text().unwrap_or_default(); + anyhow::bail!("MCP HTTP server returned {}: {}", status, text); + } + + let response: Value = resp.json() + .map_err(|e| anyhow::anyhow!("invalid JSON from MCP HTTP server: {}", e))?; + + if let Some(err) = response.get("error") { + anyhow::bail!("MCP HTTP error: {}", err); + } + + let result = response.get("result").cloned().unwrap_or(Value::Null); + extract_text_content(&result) +} + +fn extract_text_content(result: &Value) -> anyhow::Result { + if let Some(content) = result.get("content") { + if let Some(arr) = content.as_array() { + let text: Vec = arr.iter().filter_map(|item| { + if item.get("type").and_then(|t| t.as_str()) == Some("text") { + item.get("text").and_then(|t| t.as_str()).map(|s| s.to_string()) + } else { + None + } + }).collect(); + if !text.is_empty() { + return Ok(text.join("\n")); + } + } + } + Ok(serde_json::to_string_pretty(result).unwrap_or_else(|_| result.to_string())) +} + #[derive(Debug, Clone)] pub struct McpManager { pub servers: Vec, - pub running: bool, } pub struct McpToolAdapter { - name: &'static str, - description: &'static str, - parameters: serde_json::Value, + pub tool_name: String, + pub server_name: String, + pub transport: McpTransport, + pub description: String, + pub parameters: Value, } impl crate::tool::Tool for McpToolAdapter { fn name(&self) -> &'static str { - self.name + Box::leak(format!("mcp__{}__{}", self.server_name, self.tool_name).into_boxed_str()) } fn description(&self) -> &'static str { - self.description + Box::leak(self.description.clone().into_boxed_str()) } - fn parameters(&self) -> serde_json::Value { + fn parameters(&self) -> Value { self.parameters.clone() } - fn run(&self, _ctx: &crate::tool::ToolCtx, _args: &serde_json::Value) -> anyhow::Result { - Err(anyhow::anyhow!("MCP tool execution not yet implemented")) + fn run(&self, _ctx: &crate::tool::ToolCtx, args: &Value) -> anyhow::Result { + match &self.transport { + McpTransport::Stdio { command, args: extra_args } => { + call_via_stdio(command, extra_args, &self.tool_name, args) + } + McpTransport::StreamableHttp { url } => { + call_via_http(url, &self.tool_name, args) + } + } } } @@ -70,44 +243,20 @@ impl McpManager { pub fn new() -> Self { McpManager { servers: Vec::new(), - running: false, } } - pub fn add_server(&mut self, server: McpServer) { - self.servers.push(server); - } - - pub fn remove_server(&mut self, name: &str) { - self.servers.retain(|s| s.name != name); - } - - pub fn get_server(&self, name: &str) -> Option<&McpServer> { - self.servers.iter().find(|s| s.name == name) - } - - pub fn all_tools(&self) -> Vec<&McpToolInfo> { - self.servers.iter().flat_map(|s| s.tools.iter()).collect() - } - - pub fn start_all(&mut self) -> anyhow::Result<()> { - self.running = true; - Ok(()) - } - - pub fn stop_all(&mut self) -> anyhow::Result<()> { - self.running = false; - Ok(()) - } - pub fn as_tools(&self) -> Vec> { - self.all_tools().into_iter().map(|info| { - let name = format!("mcp__{}", info.name); - Box::new(McpToolAdapter { - name: Box::leak(name.into_boxed_str()), - description: Box::leak(info.description.clone().into_boxed_str()), - parameters: info.input_schema.clone(), - }) as Box + self.servers.iter().flat_map(|server| { + server.tools.iter().map(|info| { + Box::new(McpToolAdapter { + tool_name: info.name.clone(), + server_name: server.name.clone(), + transport: server.transport.clone(), + description: info.description.clone(), + parameters: info.input_schema.clone(), + }) as Box + }) }).collect() } } diff --git a/src/app/mod.rs b/src/app/mod.rs index c7f55b5..7685247 100644 --- a/src/app/mod.rs +++ b/src/app/mod.rs @@ -8,4 +8,3 @@ pub mod subagent; pub mod review; pub mod bgbash; pub mod mcp; -pub mod sec; diff --git a/src/app/mode/effort.rs b/src/app/mode/effort.rs index afdde85..a5fb57c 100644 --- a/src/app/mode/effort.rs +++ b/src/app/mode/effort.rs @@ -16,9 +16,3 @@ pub fn cycle_effort(state: &mut AppStateRest) { state.misc.effort_level = (current + 1) % EFFORT_LEVELS.len(); state.dirty = true; } - -pub fn set_effort(state: &mut AppStateRest, level: usize) { - let clamped = level.min(EFFORT_LEVELS.len() - 1); - state.misc.effort_level = clamped; - state.dirty = true; -} diff --git a/src/app/mode/mod.rs b/src/app/mode/mod.rs index 164da59..21fc21e 100644 --- a/src/app/mode/mod.rs +++ b/src/app/mode/mod.rs @@ -1,32 +1,10 @@ use serde::{Deserialize, Serialize}; -#[expect(dead_code)] -pub mod agents; pub mod bash; -#[expect(dead_code)] -pub mod editor; -#[expect(dead_code)] -pub mod effort; -#[expect(dead_code)] -pub mod help; -#[expect(dead_code)] -pub mod key_input; -#[expect(dead_code)] -pub mod loading; -#[expect(dead_code)] -pub mod mcp; -#[expect(dead_code)] -pub mod onboard; -#[expect(dead_code)] -pub mod onboard_provider; pub mod quit_confirm; -#[expect(dead_code)] -pub mod rewind; pub mod security; pub mod settings; pub mod todo; -#[expect(dead_code)] -pub mod workflow; #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum ModeKind { @@ -48,31 +26,3 @@ pub enum ModeKind { Rewind, Loading, } - -impl ModeKind { - pub fn name(&self) -> &'static str { - match self { - ModeKind::Chat => "Chat", - ModeKind::Agents => "Agents", - ModeKind::Bash => "Bash", - ModeKind::Workflow => "Workflow", - ModeKind::Help => "Help", - ModeKind::Settings => "Settings", - ModeKind::QuitConfirm => "QuitConfirm", - ModeKind::Onboard => "Onboard", - ModeKind::OnboardProvider => "OnboardProvider", - ModeKind::KeyInput => "KeyInput", - ModeKind::Editor => "Editor", - ModeKind::Effort => "Effort", - ModeKind::Mcp => "MCP", - ModeKind::Security => "Security", - ModeKind::Todo => "Todo", - ModeKind::Rewind => "Rewind", - ModeKind::Loading => "Loading", - } - } - - pub fn is_overlay(self) -> bool { - !matches!(self, ModeKind::Chat | ModeKind::Agents | ModeKind::Bash | ModeKind::Workflow) - } -} diff --git a/src/app/review/mod.rs b/src/app/review/mod.rs index 8c7884b..1c4050b 100644 --- a/src/app/review/mod.rs +++ b/src/app/review/mod.rs @@ -1,4 +1,3 @@ -use std::collections::HashMap; use std::process::Command; use crate::app::state::rest::AppStateRest; use crate::app::state::runtime::TurnEvent; @@ -50,147 +49,6 @@ pub struct Lesson { pub provenance: Provenance, } -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -pub enum ViolationEscalation { - None, - Warning, - Escalate, - Block, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -pub enum ShadowStatus { - Trial, - Graduated, - Rejected, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ShadowCheck { - pub pattern: String, - pub trial_window: u32, - pub trial_count: u32, - pub trial_passed: u32, - pub status: ShadowStatus, -} - -pub struct ReviewSystem { - pub pending: bool, - pub queue_capacity: usize, - pub repeated_violations: HashMap, - pub shadow_violations: Vec, - pub graduated_checks: Vec, - pub violation_window: u32, -} - -impl ReviewSystem { - pub fn new() -> Self { - ReviewSystem { - pending: false, - queue_capacity: 1, - repeated_violations: HashMap::new(), - shadow_violations: Vec::new(), - graduated_checks: Vec::new(), - violation_window: 10, - } - } - pub fn record_shadow_hit(&mut self, pattern: &str) -> bool { - for check in &mut self.shadow_violations { - if check.pattern == pattern { - check.trial_count += 1; - check.trial_passed += 1; - return check.trial_count >= check.trial_window; - } - } - self.shadow_violations.push(ShadowCheck { - pattern: pattern.to_string(), - trial_window: 10, - trial_count: 1, - trial_passed: 1, - status: ShadowStatus::Trial, - }); - false - } - pub fn evaluate_shadow_trials(&mut self) -> Vec { - let mut graduated = Vec::new(); - let mut remaining = Vec::new(); - for mut check in self.shadow_violations.drain(..) { - if check.trial_count < check.trial_window { - remaining.push(check); - continue; - } - let ratio = check.trial_passed as f64 / check.trial_window as f64; - let p = check.pattern.clone(); - let tp = check.trial_passed; - let tw = check.trial_window; - if ratio >= 0.3 { - self.graduated_checks.push(crate::tool::GraduatedCheck { - name: p.clone(), - pattern: p.clone(), - rule: p.clone(), - }); - graduated.push(format!("{} (graduated, fired {}/{} writes)", p, tp, tw)); - } else { - check.status = ShadowStatus::Rejected; - graduated.push(format!("{} (demoted, only {}/{} — below 30% threshold)", p, tp, tw)); - remaining.push(check); - } - } - self.shadow_violations = remaining; - graduated - } - - pub fn reset(&mut self) { - self.pending = false; - } - - pub fn check_escalation(&self, pattern: &str) -> ViolationEscalation { - let count = self.repeated_violations.get(pattern).copied().unwrap_or(0); - match count { - 0 | 1 => ViolationEscalation::None, - 2 => ViolationEscalation::Warning, - 3 | 4 => ViolationEscalation::Escalate, - _ => ViolationEscalation::Block, - } - } - - pub fn increment_violation(&mut self, pattern: &str) -> ViolationEscalation { - let entry = self.repeated_violations.entry(pattern.to_string()).or_insert(0); - *entry += 1; - self.check_escalation(pattern) - } - - pub fn should_skip_review(consecutive_empty: u32) -> bool { - consecutive_empty >= 3 - } -} - -pub fn create_pending_lesson(name: &str, content: &str, provenance: Provenance) -> Lesson { - Lesson { - name: name.to_string(), - content: content.to_string(), - confidence: Confidence::Unverified, - outcome: None, - lifecycle: LessonLifecycle::New, - scope: LessonScope::Project, - contradiction_with: None, - provenance, - } -} - -pub fn apply_lesson_calibration(state: &mut AppStateRest, lesson: &Lesson) { - if lesson.confidence != Confidence::Unverified { - return; - } - if state.mode.auto_approve() { - return; - } - state.push_toast(Toast::new( - ToastKind::Lesson, - format!("Lesson '{}' is pending review. Keep or discard?", lesson.name), - )); -} - pub fn should_trigger_review(state: &AppStateRest, origin: Origin) -> bool { if state.mode == AgentMode::Plan { return false; @@ -451,64 +309,6 @@ pub fn trigger_review(state: &mut AppStateRest) -> anyhow::Result<()> { Ok(()) } -pub fn record_review_outcome( - lessons_found: usize, - state: &mut AppStateRest, -) -> Option { - let runtime = match &mut state.session_runtime { - Some(ref mut r) => r, - None => return None, - }; - - if lessons_found > 0 { - runtime.consecutive_empty_reviews = 0; - runtime.review_count += 1; - None - } else { - runtime.consecutive_empty_reviews += 1; - runtime.review_count += 1; - None - } -} -pub fn check_violation_escalation( - pattern: &str, - system: &mut ReviewSystem, -) -> Option { - let level = system.increment_violation(pattern); - match level { - ViolationEscalation::None => None, - ViolationEscalation::Warning => Some(level), - ViolationEscalation::Escalate => Some(level), - ViolationEscalation::Block => Some(level), - } -} -pub fn format_escalation_note(pattern: &str, level: ViolationEscalation) -> String { - let label = match level { - ViolationEscalation::None => "none", - ViolationEscalation::Warning => "WARNING", - ViolationEscalation::Escalate => "ESCALATION", - ViolationEscalation::Block => "BLOCKED", - }; - format!( - "[{}] Repeated violation: '{}' has been flagged by quality review {} time(s). {}", - label, - pattern, - match level { - ViolationEscalation::None | ViolationEscalation::Warning => 2, - ViolationEscalation::Escalate => 3, - ViolationEscalation::Block => 5, - }, - match level { - ViolationEscalation::Warning => - "This pattern has appeared twice. Consider reviewing the related guideline.".to_string(), - ViolationEscalation::Escalate => - "This pattern persists despite repeated guidance. Manual review recommended.".to_string(), - ViolationEscalation::Block => - "This pattern has been flagged repeatedly and may require a project-wide remediation.".to_string(), - _ => String::new(), - } - ) -} const STALE_AFTER_DAYS: i64 = 60; @@ -545,51 +345,6 @@ pub fn maybe_run_staleness_sweep(state: &mut AppStateRest) { } } -pub fn detect_contradiction( - new_text: &str, - existing_lessons: &[crate::model::memory::Memory], -) -> Option { - let new_words: std::collections::HashSet = new_text - .to_lowercase() - .split(|c: char| !c.is_alphanumeric()) - .filter(|w| w.len() >= 4 && !is_stop_word(w)) - .map(|w| w.to_string()) - .collect(); - let opposite_markers = ["not", "never", "avoid", "don't", "do not", "instead"]; - for existing in existing_lessons { - let existing_lower = existing.content.to_lowercase(); - let exist_words: std::collections::HashSet = existing_lower - .split(|c: char| !c.is_alphanumeric()) - .filter(|w| w.len() >= 4 && !is_stop_word(w)) - .map(|w| w.to_string()) - .collect(); - - let shared = new_words.intersection(&exist_words).count(); - if shared >= 3 { - let new_has_opposite = opposite_markers.iter().any(|m| new_text.to_lowercase().contains(m)); - let old_has_opposite = opposite_markers.iter().any(|m| existing_lower.contains(m)); - if new_has_opposite != old_has_opposite { - return Some(existing.name.clone()); - } - } - } - None -} - -fn is_stop_word(w: &str) -> bool { - matches!( - w, - "this" | "that" | "with" | "from" | "have" | "been" | "were" | "they" - | "which" | "what" | "when" | "where" | "would" | "could" | "should" - | "about" | "after" | "before" | "between" | "other" | "every" | "still" | "also" - | "than" | "then" | "into" | "over" | "such" | "only" | "more" | "very" | "just" - | "because" | "while" | "being" | "made" | "make" | "does" | "done" | "using" - | "used" | "uses" | "like" | "well" | "back" | "much" | "some" | "these" | "those" - | "each" | "both" | "most" | "upon" | "here" | "down" | "your" | "its" | "our" - | "him" | "her" | "them" - ) -} - #[derive(Debug, Clone, Serialize, Deserialize)] pub struct PendingLesson { pub lesson: Lesson, @@ -611,15 +366,6 @@ pub fn save_pending_lessons(session_dir: &std::path::Path, pending: &[PendingLes std::fs::write(&path, data) } -pub fn add_pending_lesson(session_dir: &std::path::Path, lesson: Lesson, auto_resolve: bool) -> std::io::Result<()> { - let mut pending = load_pending_lessons(session_dir); - pending.push(PendingLesson { - lesson, - created_at: chrono::Utc::now().timestamp_millis(), - auto_resolve, - }); - save_pending_lessons(session_dir, &pending) -} pub fn process_pending_lessons(session_dir: &std::path::Path, memory_dir: &std::path::Path) -> std::io::Result> { let pending = load_pending_lessons(session_dir); let now = chrono::Utc::now().timestamp_millis(); diff --git a/src/app/runtime/actions/mod.rs b/src/app/runtime/actions/mod.rs index afa18d6..70c18a7 100644 --- a/src/app/runtime/actions/mod.rs +++ b/src/app/runtime/actions/mod.rs @@ -15,7 +15,6 @@ const MAX_AGENT_STEPS: usize = 40; #[derive(Debug, Clone)] pub enum Action { - Quit, ForceQuit, SwitchMode(ModeKind), SubmitInput(String), @@ -32,19 +31,10 @@ pub enum Action { CloseOverlay, ToggleYoloArm, CycleAgentMode, - ToolResult { - tool_call_id: String, - output: String, - is_error: bool, - }, - StreamToken(String), - StreamDone, - StreamError(String), SystemNote { kind: String, message: String, }, - RunCommand(String), QuitConfirm, Resize(u16, u16), Tick, @@ -60,15 +50,13 @@ pub enum Action { LessonReject { name: String, }, + StartOAuth { + provider: String, + }, } pub fn apply_action(state: &mut AppStateRest, action: Action) { match action { - Action::Quit => { - save_current_session(state); - auto_create_retrospective(state); - state.quit = true; - } Action::ForceQuit => { save_current_session(state); state.quit = true; @@ -106,8 +94,8 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) { if let Some(ref mut rt) = state.session_runtime { rt.push_message(ChatMessage::user(text)); } + state.misc.thinking = true; spawn_turn(state); - state.push_transcript(ChatMessageDisplay::new(Role::Assistant, "Thinking...".to_string())); state.dirty = true; } Action::InsertChar(c) => { @@ -137,9 +125,7 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) { state.dirty = true; } Action::ScrollUp => { - let total = state.transcript_cache.messages.len(); state.scroll.scroll_up(); - state.scroll.scroll_down(total); state.dirty = true; } Action::ScrollDown => { @@ -166,51 +152,6 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) { state.push_toast(toast); state.dirty = true; } - Action::ToolResult { - tool_call_id, - output, - is_error, - } => { - if let Some(ref mut rt) = state.session_runtime { - rt.push_message(ChatMessage::tool_result(tool_call_id.clone(), output.clone())); - rt.tool_call_results.push( - crate::app::state::runtime::ToolCallResult { - tool_call_id, - tool_name: String::new(), - output, - is_error, - duration_ms: 0, - }, - ); - } - state.dirty = true; - } - Action::StreamToken(token) => { - if let Some(ref mut rt) = state.session_runtime { - let found = rt.messages.iter_mut().rev().find(|m| { - matches!(m.role, crate::dto::chat::message::Role::Assistant) - }); - if let Some(last) = found { - let current = last.content.take().unwrap_or_default(); - last.content = Some(current + &token); - } else { - let mut msg = ChatMessage::assistant(None); - msg.content = Some(token); - rt.push_message(msg); - } - } - state.dirty = true; - } - Action::StreamDone => { - state.dirty = true; - } - Action::StreamError(msg) => { - let toast = crate::app::state::types::Toast::new( - crate::app::state::types::ToastKind::Error, - msg, - ); - state.push_toast(toast); - } Action::SystemNote { kind: _kind, message } => { let toast = crate::app::state::types::Toast::new( crate::app::state::types::ToastKind::Info, @@ -218,12 +159,6 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) { ); state.push_toast(toast); } - Action::RunCommand(text) => { - if let Some(ref mut rt) = state.session_runtime { - rt.push_message(ChatMessage::user(text)); - } - state.dirty = true; - } Action::QuitConfirm => { state.misc.overlay = Overlay::QuitConfirm; state.dirty = true; @@ -271,6 +206,26 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) { } state.dirty = true; } + Action::StartOAuth { provider } => { + let turn_events = state.turn_events.clone(); + let provider_clone = provider.clone(); + std::thread::spawn(move || { + let result = run_oauth_flow(&provider_clone); + let message = match result { + Ok(msg) => msg, + Err(e) => format!("OAuth login failed: {}", e), + }; + if let Ok(mut q) = turn_events.lock() { + q.push_back(TurnEvent::SystemNote { + kind: "oauth".to_string(), + message, + }); + } + }); + let toast = Toast::new(ToastKind::Info, format!("Opening browser for {} login...", provider)); + state.push_toast(toast); + state.dirty = true; + } Action::Tick => { let now_ms = chrono::Utc::now().timestamp_millis(); state.misc.drain_expired_toasts(now_ms); @@ -290,27 +245,17 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) { for event in events { match event { TurnEvent::AssistantMessage(msg) => { + state.misc.thinking = false; let display_content = msg.content.clone().unwrap_or_default(); if !display_content.is_empty() { - let replaced = if let Some(last) = state.transcript_cache.messages.last_mut() { - if last.role == Role::Assistant && last.content == "Thinking..." { - last.content = display_content.clone(); - true - } else { - false - } - } else { - false - }; - if !replaced { - state.push_transcript(ChatMessageDisplay::new(Role::Assistant, display_content)); - } + state.push_transcript(ChatMessageDisplay::new(Role::Assistant, display_content)); } if let Some(ref mut rt) = state.session_runtime { rt.push_message(msg); } } TurnEvent::ToolResult { tool_call_id, tool_name, output, is_error, path } => { + state.misc.thinking = false; let display_path = path.unwrap_or_default(); let display = if tool_name == "read" { let line_count = output.lines().count(); @@ -385,6 +330,7 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) { turn_finished = true; } TurnEvent::Done => { + state.misc.thinking = false; turn_finished = true; } } @@ -454,6 +400,9 @@ fn spawn_turn(state: &AppStateRest) { let events_q = turn_events.clone(); std::thread::spawn(move || { + let db = crate::model::msglog::open_or_create(&edit_session_dir) + .ok() + .map(|c| std::sync::Arc::new(std::sync::Mutex::new(c))); let tc = TurnCtx { client: crate::service::provider::LlmClient::new(api_key, model), tdefs: tool_defs, @@ -463,6 +412,7 @@ fn spawn_turn(state: &AppStateRest) { workspace_roots, edit_log_session_dir: edit_session_dir, session_id, + db, }; let result = run_agent_turn(tc, &messages, &events_q); if let Err(e) = result { @@ -485,6 +435,15 @@ struct TurnCtx { workspace_roots: Vec, edit_log_session_dir: std::path::PathBuf, session_id: String, + db: Option>>, +} + +fn archive_message(db: &Option>>, session_id: &str, msg: &ChatMessage) { + if let Some(ref arc) = db { + if let Ok(conn) = arc.lock() { + let _ = crate::model::msglog::insert_message(&conn, session_id, msg); + } + } } fn run_agent_turn( @@ -495,6 +454,7 @@ fn run_agent_turn( let mut msgs = messages.to_vec(); let mut edits_this_turn = 0u32; let mut tool_only_rounds = 0usize; + let mut prev_shaped = false; let system_text = format!( "{}\n\n{}", @@ -502,20 +462,37 @@ fn run_agent_turn( crate::resources::SYSTEM_TOOLS, ); if !msgs.iter().any(|m| matches!(m.role, crate::dto::chat::message::Role::System)) { - msgs.insert(0, ChatMessage::system(system_text)); + let sys = ChatMessage::system(system_text); + archive_message(&tc.db, &tc.session_id, &sys); + msgs.insert(0, sys); } for _step in 0..MAX_AGENT_STEPS { if tool_only_rounds >= MAX_TOOL_ONLY_TURNS { - msgs.push(ChatMessage::user( + let stop_msg = ChatMessage::user( "Stop calling tools. Respond naturally now.".to_string(), - )); + ); + archive_message(&tc.db, &tc.session_id, &stop_msg); + msgs.push(stop_msg); tool_only_rounds = 0; } + let wire_msgs = if crate::app::runtime::shortsend::should_shape(msgs.len(), prev_shaped) { + let total_chars: usize = msgs.iter() + .filter_map(|m| m.content.as_deref()) + .map(|c| c.len()) + .sum(); + let token_estimate = total_chars / 4; + prev_shaped = true; + crate::app::runtime::shortsend::shape_messages(&msgs, token_estimate) + } else { + prev_shaped = false; + msgs.clone() + }; + let response = tc .client - .chat_with_tools(&msgs, Some(tc.tdefs.clone()))?; + .chat_with_tools(&wire_msgs, Some(tc.tdefs.clone()))?; let has_tool_calls = response.tool_calls.is_some() && response.tool_calls.as_ref().is_some_and(|tc| !tc.is_empty()); @@ -524,6 +501,7 @@ fn run_agent_turn( if has_tool_calls { tool_only_rounds += 1; let tool_calls = response.tool_calls.clone().unwrap_or_default(); + archive_message(&tc.db, &tc.session_id, &response); msgs.push(response); for tool_call in tool_calls { let tool_name = tool_call.function.name.clone(); @@ -581,10 +559,12 @@ fn run_agent_turn( } let tool_msg = ChatMessage::tool_result(tool_call.id.clone(), output); + archive_message(&tc.db, &tc.session_id, &tool_msg); msgs.push(tool_msg); } } else { if !content.is_empty() { + archive_message(&tc.db, &tc.session_id, &response); if let Ok(mut q) = events_q.lock() { q.push_back(TurnEvent::AssistantMessage(response)); } @@ -698,36 +678,80 @@ fn save_current_session(state: &AppStateRest) { } } -fn auto_create_retrospective(state: &mut AppStateRest) { - if state.session_runtime.is_none() { - return; - } - let session = crate::model::session::Session::new( - state.session_id.clone(), - "session".to_string(), - ); - match crate::model::memory::auto_create_retrospective(&state.session_dir, &session) { - Ok(Some(retro)) => { - state.push_toast(crate::app::state::types::Toast::new( - crate::app::state::types::ToastKind::Info, - format!("Retrospective created: {}", retro.name), - )); - } - Ok(None) => {} - Err(e) => { - let _ = e; - } - } - let lessons: Vec = crate::model::memory::Memory::list(&state.memory_dir) - .iter() - .filter_map(|n| crate::model::memory::Memory::read(&state.memory_dir, n).ok()) - .filter(|m| m.kind == "lesson") - .collect(); - if let Some(global_dir) = dirs::data_dir().map(|d| d.join("zesdex")) { - for lesson in &lessons { - if lesson.scope.as_deref() != Some("global") { - let _ = crate::model::memory::promote_with_consensus(&global_dir, lesson); +fn run_oauth_flow(provider: &str) -> anyhow::Result { + use crate::service::oauth::manager::{OAuthConfig, OAuthManager}; + use crate::service::oauth::loopback::LoopbackServer; + use crate::service::oauth::pkce::CodeVerifier; + + let config = match provider { + "zen" | "opencode" => OAuthConfig { + auth_url: "https://opencode.ai/zen/oauth/authorize".to_string(), + token_url: "https://opencode.ai/zen/oauth/token".to_string(), + client_id: std::env::var("ZEN_CLIENT_ID") + .unwrap_or_else(|_| "zesdex".to_string()), + client_secret: std::env::var("ZEN_CLIENT_SECRET").ok(), + scopes: vec!["openid".to_string(), "profile".to_string(), "email".to_string()], + }, + "openai" => OAuthConfig { + auth_url: "https://auth0.openai.com/authorize".to_string(), + token_url: "https://auth0.openai.com/oauth/token".to_string(), + client_id: std::env::var("OPENAI_CLIENT_ID") + .unwrap_or_else(|_| "zesdex".to_string()), + client_secret: std::env::var("OPENAI_CLIENT_SECRET").ok(), + scopes: vec!["openid".to_string(), "profile".to_string(), "email".to_string()], + }, + other => { + let auth_url = std::env::var(format!("{}_AUTH_URL", other.to_uppercase())) + .map_err(|_| anyhow::anyhow!("unknown provider '{}'. Set {}_AUTH_URL env var.", other, other.to_uppercase()))?; + let token_url = std::env::var(format!("{}_TOKEN_URL", other.to_uppercase())) + .map_err(|_| anyhow::anyhow!("{}_TOKEN_URL not set", other.to_uppercase()))?; + let client_id = std::env::var(format!("{}_CLIENT_ID", other.to_uppercase())) + .unwrap_or_else(|_| "zesdex".to_string()); + OAuthConfig { + auth_url, + token_url, + client_id, + client_secret: std::env::var(format!("{}_CLIENT_SECRET", other.to_uppercase())).ok(), + scopes: vec!["openid".to_string(), "profile".to_string(), "email".to_string()], } } + }; + + let server = LoopbackServer::bind()?; + let redirect_uri = server.redirect_uri(); + + let verifier = CodeVerifier::new(); + let challenge = verifier.challenge(); + let state_token = format!("{:x}", sha2::Sha256::digest(rand_bytes(16))); + + let mut manager = OAuthManager::new(config.clone()); + let auth_url = manager.build_auth_url(&redirect_uri, &state_token, challenge.as_str()); + + let _ = webbrowser::open(&auth_url); + + let code = server.wait_for_code(120_000)?; + + manager.exchange_code(&code, &redirect_uri, verifier.as_str()) + .map_err(|e| anyhow::anyhow!("{}", e))?; + + if let Some(ref token) = manager.token { + let token_path = dirs::config_dir() + .unwrap_or_else(|| std::path::PathBuf::from(".")) + .join("zesdex") + .join(format!("oauth_{}.json", provider)); + if let Some(parent) = token_path.parent() { + let _ = std::fs::create_dir_all(parent); + } + let _ = std::fs::write(&token_path, serde_json::to_string_pretty(token).unwrap_or_default()); } + + Ok(format!("Successfully authenticated with {}.", provider)) } + +fn rand_bytes(n: usize) -> Vec { + use std::time::{SystemTime, UNIX_EPOCH}; + let seed = SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().subsec_nanos(); + (0..n).map(|i| ((seed >> (i % 4 * 8)) ^ (i as u32 * 2654435761)) as u8).collect() +} + + diff --git a/src/app/runtime/commands.rs b/src/app/runtime/commands.rs index 039235c..3e0e93b 100644 --- a/src/app/runtime/commands.rs +++ b/src/app/runtime/commands.rs @@ -41,10 +41,7 @@ pub fn apply_command(command: Command) -> Vec { }] } Command::Login { provider } => { - vec![Action::SystemNote { - kind: "oauth".to_string(), - message: format!("OAuth login flow started for {}", provider), - }] + vec![Action::StartOAuth { provider }] } Command::Unknown(cmd) => { vec![Action::SystemNote { diff --git a/src/app/runtime/event_loop/mod.rs b/src/app/runtime/event_loop/mod.rs index e69de29..89f3626 100644 --- a/src/app/runtime/event_loop/mod.rs +++ b/src/app/runtime/event_loop/mod.rs @@ -0,0 +1,52 @@ +use std::collections::VecDeque; +use std::time::{Duration, Instant}; + +use crate::app::state::runtime::TurnEvent; + +const FAST_POLL_MS: u64 = 8; +const SLOW_POLL_MS: u64 = 100; +const IDLE_THRESHOLD_MS: u64 = 500; + +pub struct EventLoop { + last_activity: Instant, + fast_poll_until: Option, +} + +impl EventLoop { + pub fn new() -> Self { + EventLoop { + last_activity: Instant::now(), + fast_poll_until: None, + } + } + + pub fn poll_interval(&self) -> Duration { + if let Some(fast_until) = self.fast_poll_until { + if Instant::now() < fast_until { + return Duration::from_millis(FAST_POLL_MS); + } + } + Duration::from_millis(SLOW_POLL_MS) + } + + pub fn mark_active(&mut self) { + self.last_activity = Instant::now(); + self.fast_poll_until = Some(Instant::now() + Duration::from_millis(IDLE_THRESHOLD_MS)); + } + + pub fn is_idle(&self) -> bool { + self.last_activity.elapsed().as_millis() as u64 > IDLE_THRESHOLD_MS + } + + pub fn drain_events( + events: &std::sync::Mutex>, + ) -> Vec { + events.lock().map(|mut q| q.drain(..).collect()).unwrap_or_default() + } +} + +impl Default for EventLoop { + fn default() -> Self { + Self::new() + } +} diff --git a/src/app/runtime/mod.rs b/src/app/runtime/mod.rs index 9fb67ca..ecb5fe3 100644 --- a/src/app/runtime/mod.rs +++ b/src/app/runtime/mod.rs @@ -1,4 +1,3 @@ pub mod actions; pub mod commands; -pub mod event_loop; -pub mod stream; +pub mod shortsend; diff --git a/src/app/runtime/stream/mod.rs b/src/app/runtime/stream/mod.rs index e69de29..03f918b 100644 --- a/src/app/runtime/stream/mod.rs +++ b/src/app/runtime/stream/mod.rs @@ -0,0 +1,172 @@ +pub mod turn; +pub mod tools; + +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum StreamEvent { + Token(String), + Reasoning(String), + ToolCallDelta { + index: usize, + id: Option, + name: Option, + arguments_delta: String, + }, + Usage { + prompt_tokens: u64, + completion_tokens: u64, + total_tokens: u64, + }, + Done, + Error(String), +} + +pub struct SseParser { + buffer: String, + event_type: Option, + data_lines: Vec, +} + +impl SseParser { + pub fn new() -> Self { + SseParser { + buffer: String::new(), + event_type: None, + data_lines: Vec::new(), + } + } + + pub fn feed(&mut self, chunk: &str) -> Vec { + self.buffer.push_str(chunk); + let mut events = Vec::new(); + while let Some(line_end) = self.buffer.find('\n') { + let line = self.buffer[..line_end].trim_end_matches('\r').to_string(); + self.buffer = self.buffer[line_end + 1..].to_string(); + if line.is_empty() { + if let Some(event) = self.flush_event() { + events.push(event); + } + } else if let Some(ty) = line.strip_prefix("event: ") { + self.event_type = Some(ty.trim().to_string()); + } else if let Some(data) = line.strip_prefix("data: ") { + self.data_lines.push(data.to_string()); + } else if line.starts_with("data:") { + self.data_lines.push(String::new()); + } + } + events + } + + fn flush_event(&mut self) -> Option { + let data = self.data_lines.join("\n"); + self.data_lines.clear(); + let event_type = self.event_type.take().unwrap_or_default(); + if data.is_empty() || data == "[DONE]" { + if data == "[DONE]" { + return Some(StreamEvent::Done); + } + return None; + } + let value: Value = serde_json::from_str(&data).ok()?; + match event_type.as_str() { + "message.stop" => Some(StreamEvent::Done), + "message.start" => None, + "message.delta" | "" => { + let delta = value.get("delta").or_else(|| value.get("choices"))?; + if let Some(choices) = delta.as_array() { + let choice = choices.first()?; + let delta = choice.get("delta")?; + if let Some(content) = delta.get("content").and_then(|c| c.as_str()) { + return Some(StreamEvent::Token(content.to_string())); + } + if let Some(reasoning) = delta.get("reasoning_content").and_then(|r| r.as_str()) { + return Some(StreamEvent::Reasoning(reasoning.to_string())); + } + if let Some(tool_calls) = delta.get("tool_calls").and_then(|tc| tc.as_array()) { + for tc in tool_calls { + let index = tc.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize; + let id = tc.get("id").and_then(|i| i.as_str()).map(|s| s.to_string()); + let name = tc.get("function") + .and_then(|f| f.get("name")) + .and_then(|n| n.as_str()) + .map(|s| s.to_string()); + let args_delta = tc.get("function") + .and_then(|f| f.get("arguments")) + .and_then(|a| a.as_str()) + .unwrap_or("") + .to_string(); + return Some(StreamEvent::ToolCallDelta { + index, + id, + name, + arguments_delta: args_delta, + }); + } + } + let finish = choice.get("finish_reason"); + if let Some(reason) = finish.and_then(|r| r.as_str()) { + if reason == "stop" || reason == "tool_calls" { + return Some(StreamEvent::Done); + } + } + } + if let Some(content) = delta.get("content").and_then(|c| c.as_str()) { + return Some(StreamEvent::Token(content.to_string())); + } + None + } + _ => None, + } + } + + pub fn reset(&mut self) { + self.buffer.clear(); + self.event_type = None; + self.data_lines.clear(); + } +} + +pub fn parse_stream_chunk(data: &str) -> Option { + let value: Value = serde_json::from_str(data).ok()?; + if value == Value::Null { + return None; + } + let choices = value.get("choices")?.as_array()?; + let choice = choices.first()?; + let delta = choice.get("delta")?; + if let Some(content) = delta.get("content").and_then(|c| c.as_str()) { + return Some(StreamEvent::Token(content.to_string())); + } + if let Some(reasoning) = delta.get("reasoning_content").and_then(|r| r.as_str()) { + return Some(StreamEvent::Reasoning(reasoning.to_string())); + } + if let Some(finish) = choice.get("finish_reason").and_then(|r| r.as_str()) { + if finish == "stop" || finish == "tool_calls" { + return Some(StreamEvent::Done); + } + } + if let Some(tool_calls) = delta.get("tool_calls").and_then(|tc| tc.as_array()) { + if let Some(tc) = tool_calls.first() { + let index = tc.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize; + let id = tc.get("id").and_then(|i| i.as_str()).map(|s| s.to_string()); + let name = tc.get("function") + .and_then(|f| f.get("name")) + .and_then(|n| n.as_str()) + .map(|s| s.to_string()); + let args = tc.get("function") + .and_then(|f| f.get("arguments")) + .and_then(|a| a.as_str()) + .unwrap_or("") + .to_string(); + return Some(StreamEvent::ToolCallDelta { + index, + id, + name, + arguments_delta: args, + }); + } + } + None +} diff --git a/src/app/runtime/stream/tools/mod.rs b/src/app/runtime/stream/tools/mod.rs index e69de29..21d294a 100644 --- a/src/app/runtime/stream/tools/mod.rs +++ b/src/app/runtime/stream/tools/mod.rs @@ -0,0 +1,73 @@ +use super::turn::ParsedToolCall; +use serde_json::{json, Value}; + +pub struct ToolCallAccumulator { + calls: Vec, +} + +impl ToolCallAccumulator { + pub fn new() -> Self { + ToolCallAccumulator { calls: Vec::new() } + } + + pub fn add_delta( + &mut self, + index: usize, + id: Option<&str>, + name: Option<&str>, + arguments_delta: &str, + ) { + while self.calls.len() <= index { + self.calls.push(ParsedToolCall { + id: String::new(), + name: String::new(), + arguments: String::new(), + is_complete: false, + }); + } + let tc = &mut self.calls[index]; + if let Some(new_id) = id { + if !new_id.is_empty() { + tc.id = new_id.to_string(); + } + } + if let Some(new_name) = name { + if !new_name.is_empty() { + tc.name = new_name.to_string(); + } + } + tc.arguments.push_str(arguments_delta); + } + + pub fn calls(&self) -> &[ParsedToolCall] { + &self.calls + } + + pub fn is_complete(&self) -> bool { + !self.calls.is_empty() && self.calls.iter().all(|tc| !tc.name.is_empty() && !tc.arguments.is_empty()) + } + + pub fn reset(&mut self) { + self.calls.clear(); + } + + pub fn pending_args(&self) -> Vec { + self.calls + .iter() + .filter(|tc| !tc.name.is_empty()) + .map(|tc| { + json!({ + "tool_call_id": tc.id, + "name": tc.name, + "arguments": tc.arguments, + }) + }) + .collect() + } +} + +impl Default for ToolCallAccumulator { + fn default() -> Self { + Self::new() + } +} diff --git a/src/app/runtime/stream/turn.rs b/src/app/runtime/stream/turn.rs index e69de29..9a7b07a 100644 --- a/src/app/runtime/stream/turn.rs +++ b/src/app/runtime/stream/turn.rs @@ -0,0 +1,131 @@ +use super::StreamEvent; +use crate::dto::chat::message::ChatMessage; +use crate::dto::chat::tool::{ToolCall, ToolFunction}; +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct StreamedTurn { + pub messages: Vec, + pub tool_calls: Vec, + pub is_complete: bool, + pub accumulated_content: String, + pub accumulated_reasoning: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ParsedToolCall { + pub id: String, + pub name: String, + pub arguments: String, + pub is_complete: bool, +} + +impl ParsedToolCall { + pub fn try_parse(&self) -> Option { + serde_json::from_str(&self.arguments).ok() + } +} + +impl StreamedTurn { + pub fn new() -> Self { + StreamedTurn { + messages: Vec::new(), + tool_calls: Vec::new(), + is_complete: false, + accumulated_content: String::new(), + accumulated_reasoning: String::new(), + } + } + + pub fn apply_event(&mut self, event: &StreamEvent) { + match event { + StreamEvent::Token(token) => { + self.accumulated_content.push_str(token); + } + StreamEvent::Reasoning(reasoning) => { + self.accumulated_reasoning.push_str(reasoning); + } + StreamEvent::ToolCallDelta { + index, + id, + name, + arguments_delta, + } => { + while self.tool_calls.len() <= *index { + self.tool_calls.push(ParsedToolCall { + id: String::new(), + name: String::new(), + arguments: String::new(), + is_complete: false, + }); + } + let tc = &mut self.tool_calls[*index]; + if let Some(new_id) = id { + if !new_id.is_empty() { + tc.id = new_id.clone(); + } + } + if let Some(new_name) = name { + if !new_name.is_empty() { + tc.name = new_name.clone(); + } + } + tc.arguments.push_str(arguments_delta); + } + StreamEvent::Done => { + self.is_complete = true; + } + _ => {} + } + } + + pub fn build_assistant_message(&self) -> ChatMessage { + let mut msg = if self.tool_calls.is_empty() { + ChatMessage::assistant(None) + } else { + let tool_dtos: Vec = self.tool_calls + .iter() + .filter(|tc| !tc.name.is_empty()) + .map(|tc| { + let args_value: serde_json::Value = serde_json::from_str(&tc.arguments) + .unwrap_or(serde_json::Value::String(tc.arguments.clone())); + ToolCall { + id: tc.id.clone(), + type_: "function".to_string(), + function: ToolFunction { + name: tc.name.clone(), + arguments: args_value, + }, + } + }) + .collect(); + let mut msg = ChatMessage::assistant(None); + if !tool_dtos.is_empty() { + msg.tool_calls = Some(tool_dtos); + } + msg + }; + let content = if self.accumulated_content.is_empty() { + None + } else { + Some(self.accumulated_content.clone()) + }; + msg.content = content; + msg + } + + pub fn has_tool_calls(&self) -> bool { + self.tool_calls.iter().any(|tc| !tc.name.is_empty()) + } + + pub fn content(&self) -> &str { + &self.accumulated_content + } +} + +impl Default for StreamedTurn { + fn default() -> Self { + Self::new() + } +} diff --git a/src/app/state/misc.rs b/src/app/state/misc.rs index 111c6b0..d1ec1c3 100644 --- a/src/app/state/misc.rs +++ b/src/app/state/misc.rs @@ -19,11 +19,6 @@ impl DirCache { let mut w = self.entries.write().await; *w = paths; } - - pub async fn get(&self) -> Vec { - let r = self.entries.read().await; - r.clone() - } } #[derive(Debug, Clone)] @@ -53,10 +48,6 @@ impl ScrollState { } } - pub fn scroll_to_bottom(&mut self, total: usize) { - self.offset = total.saturating_sub(self.max_visible); - } - pub fn set_max_visible(&mut self, max: usize) { self.max_visible = max; } @@ -202,26 +193,17 @@ impl InputState { None => {} } } - - pub fn clear(&mut self) { - self.buffer.clear(); - self.cursor = 0; - self.history_idx = None; - } } #[derive(Debug, Clone)] pub struct MiscState { pub overlay: Overlay, pub toasts: Vec, - pub dirty: bool, pub yolo_armed: bool, pub security_armed: bool, - pub security_acknowledged: bool, pub esc_press_count: u32, pub last_staleness_sweep_ms: i64, - pub effort_level: usize, - pub editor: Option, + pub thinking: bool, } impl MiscState { @@ -229,14 +211,11 @@ impl MiscState { MiscState { overlay: Overlay::None, toasts: Vec::new(), - dirty: true, yolo_armed: false, security_armed: false, - security_acknowledged: false, esc_press_count: 0, last_staleness_sweep_ms: 0, - effort_level: 1, - editor: None, + thinking: false, } } diff --git a/src/app/state/mod.rs b/src/app/state/mod.rs index e373597..094c8c6 100644 --- a/src/app/state/mod.rs +++ b/src/app/state/mod.rs @@ -1,6 +1,4 @@ pub mod misc; pub mod rest; pub mod runtime; -pub mod diff; -pub mod snapshot; pub mod types; diff --git a/src/app/state/rest.rs b/src/app/state/rest.rs index 6c0c031..ae357d6 100644 --- a/src/app/state/rest.rs +++ b/src/app/state/rest.rs @@ -28,14 +28,6 @@ impl ChatMessageDisplay { } } -#[derive(Clone)] -pub struct CronJob { - pub id: String, - pub description: String, - pub cron_expr: String, - pub active: bool, -} - #[derive(Clone)] pub struct AppStateRest { pub mode: AgentMode, @@ -46,12 +38,10 @@ pub struct AppStateRest { pub memory_dir: PathBuf, pub download_dir: PathBuf, pub worktrees_dir: PathBuf, - pub current_dir: PathBuf, pub dir_cache: Arc>, pub edit_log: EditLog, pub session_runtime: Option, pub sessions: Vec, - pub crons: Vec, pub transcript_cache: TranscriptCache, pub scroll: ScrollState, pub input: InputState, @@ -83,7 +73,6 @@ impl AppStateRest { memory_dir, download_dir, worktrees_dir, - current_dir: std::env::current_dir().unwrap_or_default(), turn_events: Arc::new(Mutex::new(VecDeque::new())), turn_in_flight: Arc::new(Mutex::new(false)), dir_cache: Arc::new(RwLock::new(dir_cache)), @@ -92,7 +81,6 @@ impl AppStateRest { workflow_engine: WorkflowEngine::new(), mcp_manager: McpManager::new(), sessions: Vec::new(), - crons: Vec::new(), transcript_cache: TranscriptCache::new(200), scroll: ScrollState::new(), input: InputState::new(), @@ -102,10 +90,6 @@ impl AppStateRest { } } - pub fn mode(&self) -> AgentMode { - self.mode - } - pub fn turn_in_flight(&self) -> bool { self.turn_in_flight.lock().map(|g| *g).unwrap_or(false) } @@ -147,7 +131,7 @@ impl AppStateRest { workspaces: self.workspace_roots.clone(), session_dir: self.session_dir.clone(), memory_dir: self.memory_dir.clone(), - download_dir: self.download_dir.clone(), + _download_dir: self.download_dir.clone(), worktrees_dir: self.worktrees_dir.clone(), dir_cache: self.dir_cache.clone(), internet_mode: self.settings.internet_mode.clone(), diff --git a/src/app/state/runtime.rs b/src/app/state/runtime.rs index 0078e15..dcfbfa3 100644 --- a/src/app/state/runtime.rs +++ b/src/app/state/runtime.rs @@ -109,15 +109,4 @@ impl SessionRuntime { pub fn push_message(&mut self, msg: crate::dto::chat::message::ChatMessage) { self.messages.push(msg); } - - pub fn record_api_call(&mut self, tokens_in: u64, tokens_out: u64, duration_ms: u64) { - self.usage.tokens_in += tokens_in; - self.usage.tokens_out += tokens_out; - self.usage.api_calls += 1; - self.usage.total_ms += duration_ms; - } - - pub fn record_review_tokens(&mut self, tokens: u64) { - self.usage.review_tokens += tokens; - } } diff --git a/src/app/state/types.rs b/src/app/state/types.rs index f3f08df..b5062e1 100644 --- a/src/app/state/types.rs +++ b/src/app/state/types.rs @@ -32,27 +32,6 @@ impl AgentMode { } } -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -pub enum PanelKind { - Chat, - Agents, - Bash, - Workflow, - Help, -} - -impl PanelKind { - pub fn name(&self) -> &'static str { - match self { - PanelKind::Chat => "Chat", - PanelKind::Agents => "Agents", - PanelKind::Bash => "Bash", - PanelKind::Workflow => "Workflow", - PanelKind::Help => "Help", - } - } -} - #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum ToastKind { Info, diff --git a/src/app/subagent/context.rs b/src/app/subagent/context.rs index fed0e8c..f88a432 100644 --- a/src/app/subagent/context.rs +++ b/src/app/subagent/context.rs @@ -4,12 +4,10 @@ use super::spawn::AgentDefinition; pub const REVIEWER_ALLOWED: &[&str] = &["read", "grep", "glob", "recall", "remember"]; pub struct SubagentContext { - pub definition: AgentDefinition, pub system_prompt: String, pub allowed_tools: Vec, pub max_steps: usize, pub session_dir: PathBuf, - pub origin: crate::app::state::types::Origin, } pub fn build_subagent_context(def: AgentDefinition) -> SubagentContext { @@ -21,11 +19,9 @@ pub fn build_subagent_context(def: AgentDefinition) -> SubagentContext { } }); SubagentContext { - definition: def, system_prompt: String::new(), allowed_tools, max_steps: 25, session_dir: PathBuf::new(), - origin: crate::app::state::types::Origin::SubAgent, } } diff --git a/src/app/subagent/engine.rs b/src/app/subagent/engine.rs index ba5ea50..e951d58 100644 --- a/src/app/subagent/engine.rs +++ b/src/app/subagent/engine.rs @@ -37,16 +37,16 @@ pub fn run_subagent(ctx: SubagentContext, tx: mpsc::Sender) -> an Ok(r) => r, Err(e) => { let _ = tx.blocking_send(SubagentEvent::StepFailed { - step, - error: e.to_string(), + _step: step, + _error: e.to_string(), }); anyhow::bail!("subagent call failed at step {}: {}", step, e); } }; let _ = tx.blocking_send(SubagentEvent::ToolCall { - tool: "api".to_string(), - args: serde_json::json!({"response": response}), + _tool: "api".to_string(), + _args: serde_json::json!({"response": response}), }); let tool_calls = tool_call_from_response(&response); @@ -54,8 +54,8 @@ pub fn run_subagent(ctx: SubagentContext, tx: mpsc::Sender) -> an output.push_str(&response); output.push('\n'); let _ = tx.blocking_send(SubagentEvent::StepCompleted { - step, - output: response.clone(), + _step: step, + _output: response.clone(), }); if !response.contains("Tool:") { break; @@ -70,8 +70,8 @@ pub fn run_subagent(ctx: SubagentContext, tx: mpsc::Sender) -> an let msg = format!("tool '{}' not allowed for this subagent", tool_name); messages.push(ChatMessage::tool_result(tool_name.clone(), msg.clone())); let _ = tx.blocking_send(SubagentEvent::ToolResult { - tool: tool_name.clone(), - output: msg, + _tool: tool_name.clone(), + _output: msg, }); continue; } @@ -80,8 +80,8 @@ pub fn run_subagent(ctx: SubagentContext, tx: mpsc::Sender) -> an let msg = format!("risky tool '{}' requires explicit permission; not allowed for this subagent", tool_name); messages.push(ChatMessage::tool_result(tool_name.clone(), msg.clone())); let _ = tx.blocking_send(SubagentEvent::ToolResult { - tool: tool_name.clone(), - output: msg, + _tool: tool_name.clone(), + _output: msg, }); continue; } @@ -94,22 +94,22 @@ pub fn run_subagent(ctx: SubagentContext, tx: mpsc::Sender) -> an match result { Ok(output_text) => { let _ = tx.blocking_send(SubagentEvent::ToolResult { - tool: tool_name.clone(), - output: output_text, + _tool: tool_name.clone(), + _output: output_text, }); } Err(e) => { let msg = format!("tool '{}' failed: {}", tool_name, e); let _ = tx.blocking_send(SubagentEvent::ToolResult { - tool: tool_name.clone(), - output: msg, + _tool: tool_name.clone(), + _output: msg, }); } } } let _ = tx.blocking_send(SubagentEvent::StepCompleted { - step, - output: response.clone(), + _step: step, + _output: response.clone(), }); } @@ -119,6 +119,6 @@ pub fn run_subagent(ctx: SubagentContext, tx: mpsc::Sender) -> an messages.push(user_msg); } - let _ = tx.blocking_send(SubagentEvent::Completed { output: output.clone() }); + let _ = tx.blocking_send(SubagentEvent::Completed { _output: output.clone() }); Ok(output) } diff --git a/src/app/subagent/event.rs b/src/app/subagent/event.rs index 19214ae..3efb9a1 100644 --- a/src/app/subagent/event.rs +++ b/src/app/subagent/event.rs @@ -3,25 +3,22 @@ use serde_json::Value; #[derive(Debug, Clone)] pub enum SubagentEvent { StepCompleted { - step: usize, - output: String, + _step: usize, + _output: String, }, StepFailed { - step: usize, - error: String, + _step: usize, + _error: String, }, Completed { - output: String, - }, - Failed { - error: String, + _output: String, }, ToolCall { - tool: String, - args: Value, + _tool: String, + _args: Value, }, ToolResult { - tool: String, - output: String, + _tool: String, + _output: String, }, } diff --git a/src/app/subagent/spawn.rs b/src/app/subagent/spawn.rs index 203f98f..345ba98 100644 --- a/src/app/subagent/spawn.rs +++ b/src/app/subagent/spawn.rs @@ -22,34 +22,8 @@ impl AgentDefinition { } } - pub fn with_system_prompt(mut self, prompt: String) -> Self { - self.system_prompt = Some(prompt); - self - } - - pub fn with_allowed_tools(mut self, tools: Vec) -> Self { - self.allowed_tools = Some(tools); - self - } - pub fn with_max_steps(mut self, steps: usize) -> Self { self.max_steps = Some(steps); self } - - pub fn with_temperature(mut self, temp: f32) -> Self { - self.temperature = Some(temp); - self - } -} - -pub fn merge_agent_defs(base: AgentDefinition, overrides: AgentDefinition) -> AgentDefinition { - AgentDefinition { - name: base.name, - role: base.role, - system_prompt: overrides.system_prompt.or(base.system_prompt), - allowed_tools: overrides.allowed_tools.or(base.allowed_tools), - max_steps: overrides.max_steps.or(base.max_steps), - temperature: overrides.temperature.or(base.temperature), - } } diff --git a/src/app/workflow/engine.rs b/src/app/workflow/engine.rs index 5557f75..d9fb201 100644 --- a/src/app/workflow/engine.rs +++ b/src/app/workflow/engine.rs @@ -1,8 +1,9 @@ use std::collections::HashMap; +use std::sync::{Arc, Mutex}; use serde::{Deserialize, Serialize}; use super::script::{ScriptPrimitive, WorkflowScript}; -static FINDINGS: std::sync::Mutex> = std::sync::Mutex::new(Vec::new()); +static FINDINGS: Mutex> = Mutex::new(Vec::new()); #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum AgentState { @@ -20,17 +21,6 @@ pub struct AgentStatus { pub error: Option, } -impl AgentStatus { - pub fn new() -> Self { - AgentStatus { - state: AgentState::Idle, - started_at: None, - completed_at: None, - error: None, - } - } -} - #[derive(Debug, Clone, Serialize, Deserialize)] pub struct WorkflowAgent { pub id: String, @@ -38,20 +28,9 @@ pub struct WorkflowAgent { pub status: AgentStatus, } -impl WorkflowAgent { - pub fn new(id: String, name: String) -> Self { - WorkflowAgent { - id, - name, - status: AgentStatus::new(), - } - } -} - #[derive(Debug, Clone)] pub struct WorkflowEngine { pub agents: Vec, - pub concurrency_cap: usize, pub findings: Vec, } @@ -59,55 +38,167 @@ impl WorkflowEngine { pub fn new() -> Self { WorkflowEngine { agents: Vec::new(), - concurrency_cap: 5, findings: Vec::new(), } } - - pub fn with_concurrency_cap(mut self, cap: usize) -> Self { - self.concurrency_cap = cap; - self - } - - pub fn add_agent(&mut self, agent: WorkflowAgent) { - self.agents.push(agent); - } } -pub fn execute_primitive(primitive: &ScriptPrimitive, args: &HashMap) -> anyhow::Result<()> { +fn spawn_single_agent(prompt: &str, findings_snapshot: Vec) -> anyhow::Result { + use crate::app::subagent::context::build_subagent_context; + use crate::app::subagent::engine::run_subagent; + use crate::app::subagent::spawn::AgentDefinition; + + let def = AgentDefinition::new("workflow-agent".to_string(), "coder".to_string()) + .with_max_steps(20); + let mut ctx = build_subagent_context(def); + + let findings_section = if findings_snapshot.is_empty() { + String::new() + } else { + format!( + "\n\nFindings from sibling agents in this workflow run:\n{}", + findings_snapshot + .iter() + .enumerate() + .map(|(i, f)| format!("{}. {}", i + 1, f)) + .collect::>() + .join("\n") + ) + }; + + ctx.system_prompt = format!("{}{}", prompt, findings_section); + + let (tx, _rx) = tokio::sync::mpsc::channel(32); + run_subagent(ctx, tx) +} + +type ParallelResult = (usize, anyhow::Result>); + +pub fn execute_primitive( + primitive: &ScriptPrimitive, + args: &HashMap, + concurrency_cap: usize, +) -> anyhow::Result> { match primitive { - ScriptPrimitive::Agent(name) => { - let _agent_name = name; - let _args = args; - Ok(()) + ScriptPrimitive::Agent(prompt) => { + let resolved = resolve_template(prompt, args); + let findings_snapshot = FINDINGS.lock().map(|f| f.clone()).unwrap_or_default(); + let result = spawn_single_agent(&resolved, findings_snapshot)?; + Ok(vec![result]) } + ScriptPrimitive::Parallel(scripts) => { - for script in scripts { - execute_primitive(script, args)?; + let semaphore = Arc::new(Semaphore::new(concurrency_cap.max(1))); + let results: Arc>> = + Arc::new(Mutex::new(Vec::new())); + + let handles: Vec<_> = scripts + .iter() + .enumerate() + .map(|(idx, script)| { + let script = script.clone(); + let args = args.clone(); + let sem = Arc::clone(&semaphore); + let results = Arc::clone(&results); + let cap = concurrency_cap; + + std::thread::spawn(move || { + let _permit = sem.acquire(); + let result = execute_primitive(&script, &args, cap); + if let Ok(mut locked) = results.lock() { + locked.push((idx, result)); + } + }) + }) + .collect(); + + for handle in handles { + let _ = handle.join(); } - Ok(()) + + let mut locked = results.lock().map_err(|_| anyhow::anyhow!("parallel results lock poisoned"))?; + locked.sort_by_key(|(idx, _)| *idx); + let mut all = Vec::new(); + for (_, res) in locked.drain(..) { + match res { + Ok(outputs) => all.extend(outputs), + Err(e) => all.push(format!("agent error: {}", e)), + } + } + Ok(all) } + ScriptPrimitive::Pipeline(scripts) => { - for script in scripts { - execute_primitive(script, args)?; + let results_store: Arc>>>> = + Arc::new(Mutex::new(vec![None; scripts.len()])); + let args_arc = Arc::new(args.clone()); + + let handles: Vec<_> = scripts + .iter() + .enumerate() + .map(|(idx, script)| { + let script = script.clone(); + let args = Arc::clone(&args_arc); + let store = Arc::clone(&results_store); + let cap = concurrency_cap; + + std::thread::spawn(move || { + let result = execute_primitive(&script, &args, cap); + if let Ok(mut locked) = store.lock() { + locked[idx] = Some(result.unwrap_or_else(|e| vec![format!("pipeline stage {} error: {}", idx, e)])); + } + }) + }) + .collect(); + + for handle in handles { + let _ = handle.join(); } - Ok(()) + + let locked = results_store.lock().map_err(|_| anyhow::anyhow!("pipeline results lock poisoned"))?; + let mut all = Vec::new(); + for outputs in locked.iter().flatten() { + all.extend(outputs.iter().cloned()); + } + Ok(all) } + ScriptPrimitive::Phase { name: _name, script } => { - execute_primitive(script, args) + execute_primitive(script, args, concurrency_cap) } } } pub fn run_workflow(script: &WorkflowScript, args: &HashMap) -> anyhow::Result { + if let Ok(mut findings) = FINDINGS.lock() { + findings.clear(); + } + let concurrency_cap = if script.options.max_concurrency > 0 { script.options.max_concurrency.min(5) } else { 5 }; - let _cap = concurrency_cap; - execute_primitive(&script.script, args)?; - Ok("workflow completed".to_string()) + + let results = execute_primitive(&script.script, args, concurrency_cap)?; + + let summary = if results.is_empty() { + "workflow completed with no output".to_string() + } else { + format!( + "workflow '{}' completed. {} agent result(s):\n{}", + script.name, + results.len(), + results + .iter() + .enumerate() + .map(|(i, r)| format!("[{}] {}", i + 1, r.lines().next().unwrap_or(r))) + .collect::>() + .join("\n") + ) + }; + + Ok(summary) } pub fn note_finding(text: &str) { @@ -115,3 +206,46 @@ pub fn note_finding(text: &str) { findings.push(text.to_string()); } } + +fn resolve_template(template: &str, args: &HashMap) -> String { + let mut result = template.to_string(); + for (key, value) in args { + result = result.replace(&format!("{{{{{}}}}}", key), value); + } + result +} + +struct Semaphore { + count: Mutex, + condvar: std::sync::Condvar, +} + +impl Semaphore { + fn new(count: usize) -> Self { + Semaphore { + count: Mutex::new(count), + condvar: std::sync::Condvar::new(), + } + } + + fn acquire(&self) -> SemaphoreGuard<'_> { + let mut count = self.count.lock().unwrap(); + while *count == 0 { + count = self.condvar.wait(count).unwrap(); + } + *count -= 1; + SemaphoreGuard { sem: self } + } +} + +struct SemaphoreGuard<'a> { + sem: &'a Semaphore, +} + +impl<'a> Drop for SemaphoreGuard<'a> { + fn drop(&mut self) { + let mut count = self.sem.count.lock().unwrap(); + *count += 1; + self.sem.condvar.notify_one(); + } +} diff --git a/src/dto/chat/message.rs b/src/dto/chat/message.rs index c82e4f6..afa65c4 100644 --- a/src/dto/chat/message.rs +++ b/src/dto/chat/message.rs @@ -13,17 +13,6 @@ pub enum Role { } impl Role { - pub fn is_user(&self) -> bool { - matches!(self, Role::User) - } - - pub fn is_assistant(&self) -> bool { - matches!(self, Role::Assistant) - } - - pub fn is_system(&self) -> bool { - matches!(self, Role::System) - } } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -79,10 +68,3 @@ impl ChatMessage { } } } - -#[derive(Debug, Clone)] -pub struct ChatMessageDisplay { - pub role: Role, - pub content: String, - pub timestamp: i64, -} diff --git a/src/dto/chat/tool.rs b/src/dto/chat/tool.rs index 4ca907a..4b69180 100644 --- a/src/dto/chat/tool.rs +++ b/src/dto/chat/tool.rs @@ -15,14 +15,6 @@ pub struct ToolFunction { pub arguments: Value, } -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ToolResult { - pub tool_call_id: String, - pub output: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub is_error: Option, -} - pub fn sanitize_tool_arguments(args: &Value) -> Value { match args { Value::String(s) => { diff --git a/src/dto/provider/response.rs b/src/dto/provider/response.rs index 881dc05..b61bb10 100644 --- a/src/dto/provider/response.rs +++ b/src/dto/provider/response.rs @@ -15,43 +15,3 @@ pub struct Choice { pub message: super::super::chat::message::ChatMessage, pub finish_reason: Option, } - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StreamChunk { - pub id: Option, - pub model: Option, - pub choices: Vec, - pub usage: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StreamChoice { - pub index: u32, - pub delta: Delta, - pub finish_reason: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct Delta { - #[serde(skip_serializing_if = "Option::is_none")] - pub role: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_calls: Option>, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DeltaToolCall { - pub index: u32, - pub id: Option, - #[serde(rename = "type")] - pub type_: Option, - pub function: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DeltaFunction { - pub name: Option, - pub arguments: Option, -} diff --git a/src/dto/provider/usage.rs b/src/dto/provider/usage.rs index df47c16..5714f25 100644 --- a/src/dto/provider/usage.rs +++ b/src/dto/provider/usage.rs @@ -10,9 +10,3 @@ pub struct Usage { #[serde(skip_serializing_if = "Option::is_none")] pub completion_tokens_cost: Option, } - -impl Usage { - pub fn total(&self) -> u32 { - self.total_tokens.unwrap_or(0) - } -} diff --git a/src/ipc/client.rs b/src/ipc/client.rs index 2e72e37..669fbf2 100644 --- a/src/ipc/client.rs +++ b/src/ipc/client.rs @@ -6,11 +6,6 @@ pub struct IpcClient { } impl IpcClient { - pub fn connect_tcp(addr: &str) -> Result { - let conn = Connection::connect_tcp(addr)?; - Ok(IpcClient { conn }) - } - pub fn connect_unix(path: &str) -> Result { let conn = Connection::connect_unix(path)?; Ok(IpcClient { conn }) @@ -23,15 +18,4 @@ impl IpcClient { pub fn receive(&mut self) -> Result> { self.conn.receive() } - - pub fn request( - &mut self, - request: &T, - ) -> Result { - self.conn.send(request)?; - match self.conn.receive::()? { - Some(response) => Ok(response), - None => anyhow::bail!("connection closed before response"), - } - } } diff --git a/src/ipc/conn.rs b/src/ipc/conn.rs index 09c981f..d7498bb 100644 --- a/src/ipc/conn.rs +++ b/src/ipc/conn.rs @@ -1,38 +1,28 @@ -use std::net::TcpStream; use std::os::unix::net::UnixStream; use anyhow::Result; use super::frame; -pub enum Connection { - Tcp(TcpStream), - Unix(UnixStream), +pub struct Connection { + inner: UnixStream, } impl Connection { - pub fn connect_tcp(addr: &str) -> Result { - let stream = TcpStream::connect(addr)?; - stream.set_nodelay(true)?; - Ok(Connection::Tcp(stream)) + pub fn from_stream(stream: UnixStream) -> Result { + Ok(Connection { inner: stream }) } pub fn connect_unix(path: &str) -> Result { let stream = UnixStream::connect(path)?; - Ok(Connection::Unix(stream)) + Ok(Connection { inner: stream }) } pub fn send(&mut self, value: &T) -> Result<()> { let data = frame::serialize_frame(value)?; - match self { - Connection::Tcp(ref mut s) => frame::write_frame(s, &data), - Connection::Unix(ref mut s) => frame::write_frame(s, &data), - } + frame::write_frame(&mut self.inner, &data) } pub fn receive(&mut self) -> Result> { - let data = match self { - Connection::Tcp(ref mut s) => frame::read_frame(s)?, - Connection::Unix(ref mut s) => frame::read_frame(s)?, - }; + let data = frame::read_frame(&mut self.inner)?; match data { Some(bytes) => { let value: T = frame::deserialize_frame(&bytes)?; @@ -41,17 +31,4 @@ impl Connection { None => Ok(None), } } - - pub fn try_clone(&self) -> Result { - match self { - Connection::Tcp(s) => { - let cloned = s.try_clone()?; - Ok(Connection::Tcp(cloned)) - } - Connection::Unix(s) => { - let cloned = s.try_clone()?; - Ok(Connection::Unix(cloned)) - } - } - } } diff --git a/src/ipc/mod.rs b/src/ipc/mod.rs index da9e460..9f3ad2e 100644 --- a/src/ipc/mod.rs +++ b/src/ipc/mod.rs @@ -1,7 +1,5 @@ pub mod client; pub mod conn; -pub mod diff; pub mod frame; pub mod protocol; pub mod server; -pub mod snapshot; diff --git a/src/ipc/server.rs b/src/ipc/server.rs index 0ad19e3..735b04d 100644 --- a/src/ipc/server.rs +++ b/src/ipc/server.rs @@ -1,84 +1,20 @@ -use std::net::TcpListener; use std::os::unix::net::UnixListener; -use std::thread; use anyhow::Result; use super::conn::Connection; -enum ListenerKind { - Tcp(TcpListener), - Unix(UnixListener), -} - pub struct IpcServer { - listener: ListenerKind, + listener: UnixListener, } impl IpcServer { - pub fn bind(addr: &str) -> Result { - let listener = TcpListener::bind(addr)?; - Ok(IpcServer { listener: ListenerKind::Tcp(listener) }) - } - pub fn bind_unix(path: &str) -> Result { let _ = std::fs::remove_file(path); let listener = UnixListener::bind(path)?; - Ok(IpcServer { listener: ListenerKind::Unix(listener) }) + Ok(IpcServer { listener }) } pub fn accept(&self) -> Result { - match &self.listener { - ListenerKind::Tcp(l) => { - let (stream, _addr) = l.accept()?; - stream.set_nodelay(true)?; - Ok(Connection::Tcp(stream)) - } - ListenerKind::Unix(l) => { - let (stream, _addr) = l.accept()?; - Ok(Connection::Unix(stream)) - } - } - } - - pub fn accept_with_handler(self, handler: F) -> thread::JoinHandle<()> - where - F: Fn(Connection) -> Result<()> + Send + 'static, - { - match self.listener { - ListenerKind::Tcp(l) => { - thread::spawn(move || { - for stream in l.incoming() { - match stream { - Ok(s) => { - let _ = s.set_nodelay(true); - if let Err(e) = handler(Connection::Tcp(s)) { - eprintln!("ipc handler error: {}", e); - } - } - Err(e) => { - eprintln!("ipc accept error: {}", e); - break; - } - } - } - }) - } - ListenerKind::Unix(l) => { - thread::spawn(move || { - for stream in l.incoming() { - match stream { - Ok(s) => { - if let Err(e) = handler(Connection::Unix(s)) { - eprintln!("ipc handler error: {}", e); - } - } - Err(e) => { - eprintln!("ipc accept error: {}", e); - break; - } - } - } - }) - } - } + let (stream, _addr) = self.listener.accept()?; + Connection::from_stream(stream) } } diff --git a/src/main.rs b/src/main.rs index 3b12c33..f31037b 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,5 +1,3 @@ -#![expect(dead_code)] - use std::io; use std::io::Write; use anyhow::Result; @@ -13,7 +11,6 @@ mod controller; mod dto; mod ipc; mod model; -mod security; mod service; mod tool; mod resources; diff --git a/src/model/editlog.rs b/src/model/editlog.rs index b31c47d..48a8b47 100644 --- a/src/model/editlog.rs +++ b/src/model/editlog.rs @@ -40,24 +40,6 @@ impl EditLog { Ok(()) } - pub fn load(path: &std::path::Path) -> std::io::Result { - let content = std::fs::read_to_string(path)?; - let entries: Vec = content - .lines() - .filter_map(|l| serde_json::from_str(l).ok()) - .collect(); - Ok(EditLog { - entries, - path: path.to_path_buf(), - }) - } - - pub fn recent(&self, n: usize) -> &[EditLogEntry] { - let len = self.entries.len(); - let start = len.saturating_sub(n); - &self.entries[start..] - } - pub fn len(&self) -> usize { self.entries.len() } @@ -73,7 +55,6 @@ mod tests { let _ = std::fs::create_dir_all(&dir); let log = EditLog::new(&dir); assert_eq!(log.len(), 0); - assert_eq!(log.recent(5).len(), 0); let _ = std::fs::remove_dir_all(&dir); } @@ -94,19 +75,10 @@ mod tests { }; log.append(entry.clone()).unwrap(); assert_eq!(log.len(), 1); - - let loaded = EditLog::load(&log.path).unwrap(); - assert_eq!(loaded.len(), 1); - assert_eq!(loaded.entries[0].reason, "test reason"); - assert_eq!(loaded.entries[0].tool, "write"); - assert_eq!(loaded.entries[0].path, "test.txt"); - - let recent = log.recent(1); - assert_eq!(recent.len(), 1); - assert_eq!(recent[0].bytes_delta, 42); - - let empty = log.recent(0); - assert_eq!(empty.len(), 0); + assert_eq!(log.entries[0].reason, "test reason"); + assert_eq!(log.entries[0].tool, "write"); + assert_eq!(log.entries[0].path, "test.txt"); + assert_eq!(log.entries[0].bytes_delta, 42); let _ = std::fs::remove_dir_all(&dir); } @@ -128,10 +100,8 @@ mod tests { }).unwrap(); } assert_eq!(log.len(), 5); - let recent = log.recent(3); - assert_eq!(recent.len(), 3); - assert_eq!(recent[0].reason, "reason 2"); - assert_eq!(recent[2].reason, "reason 4"); + assert_eq!(log.entries[0].reason, "reason 0"); + assert_eq!(log.entries[4].reason, "reason 4"); let _ = std::fs::remove_dir_all(&dir); } } diff --git a/src/model/memory.rs b/src/model/memory.rs index 2627c41..d9d5e2c 100644 --- a/src/model/memory.rs +++ b/src/model/memory.rs @@ -1,8 +1,6 @@ use std::path::{Path, PathBuf}; use serde::{Deserialize, Serialize}; -use super::session::Session; - #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Memory { pub name: String, @@ -129,16 +127,6 @@ impl Memory { }) .collect() } - - pub fn load_index(memory_dir: &Path) -> Vec { - let index_path = memory_dir.join("MEMORY.md"); - let content = std::fs::read_to_string(index_path).unwrap_or_default(); - content.lines().filter_map(|l| { - let l = l.trim(); - if l.is_empty() || l.starts_with('#') { return None; } - l.split(']').next().and_then(|s| s.split('[').nth(1)).map(|s| s.to_string()) - }).collect() - } } pub fn slug_path(memory_dir: &Path, raw: &str) -> PathBuf { @@ -159,26 +147,6 @@ pub fn export_lessons(memory_dir: &Path, output: &Path) -> std::io::Result<()> { std::fs::write(output, data)?; Ok(()) } -pub fn promote_with_consensus(global_dir: &Path, lesson: &Memory) -> std::io::Result { - let global_path = global_dir.join("memory"); - std::fs::create_dir_all(&global_path)?; - let existing = Memory::list(&global_path); - let slug = Memory::slugify(&lesson.name).unwrap_or_default(); - if existing.contains(&slug) { - return Ok(true); - } - let consensus = lesson.outcome.as_deref() == Some("verified"); - - if consensus { - let mut promoted = lesson.clone(); - promoted.scope = Some("global".to_string()); - promoted.write(&global_path)?; - Ok(true) - } else { - Ok(false) - } -} - pub fn import_lessons(memory_dir: &Path, input: &Path) -> std::io::Result { let data = std::fs::read_to_string(input)?; let lessons: Vec = serde_json::from_str(&data) @@ -194,30 +162,6 @@ pub fn import_lessons(memory_dir: &Path, input: &Path) -> std::io::Result } Ok(imported) } -pub fn auto_create_retrospective(session_dir: &Path, session: &Session) -> std::io::Result> { - let now = chrono::Utc::now().timestamp_millis(); - let session_age_ms = now.saturating_sub(session.created_at); - if session_age_ms < 60_000 { - return Ok(None); - } - let retro_name = format!("retrospective-{}", session.id); - let retro_path = Memory::path(session_dir, &retro_name); - if retro_path.exists() { - return Ok(None); - } - let lessons: Vec = Memory::list(session_dir) - .iter() - .filter_map(|n| Memory::read(session_dir, n).ok()) - .filter(|m| m.kind == "lesson") - .collect(); - - if lessons.is_empty() { - return Ok(None); - } - - let retrospective = create_retrospective(session_dir, session, &lessons)?; - Ok(Some(retrospective)) -} #[cfg(test)] mod tests { @@ -386,31 +330,3 @@ mod tests { let _ = std::fs::remove_file(&export_path); } } - -pub fn create_retrospective(session_dir: &Path, session: &Session, lessons: &[Memory]) -> std::io::Result { - let now = chrono::Utc::now().timestamp_millis(); - let lessons_content: String = lessons.iter() - .map(|l| format!("- {}: {}", l.name, l.description)) - .collect::>() - .join("\n"); - let content = format!( - "# Session Retrospective\n\nSession: {}\nCreated: {}\nLessons learned:\n{}\n", - session.title, now, lessons_content, - ); - let memory = Memory { - name: format!("retrospective-{}", session.id), - description: format!("End-of-session retrospective for {}", session.title), - content, - kind: "retrospective".to_string(), - created_at: now, - updated_at: now, - outcome: None, - lifecycle: "new".to_string(), - scope: Some("project".to_string()), - before_snippet: None, - after_snippet: None, - provenances: vec![], - }; - memory.write(session_dir)?; - Ok(memory) -} diff --git a/src/model/mod.rs b/src/model/mod.rs index a1d7062..5028989 100644 --- a/src/model/mod.rs +++ b/src/model/mod.rs @@ -1,10 +1,6 @@ -pub mod agent_def; -pub mod app_config; -pub mod conversation; pub mod editlog; pub mod memory; pub mod msglog; pub mod session; -pub mod session_lock; pub mod settings; pub mod store; diff --git a/src/model/msglog/mod.rs b/src/model/msglog/mod.rs index d696cc1..1b72bd7 100644 --- a/src/model/msglog/mod.rs +++ b/src/model/msglog/mod.rs @@ -1,4 +1,14 @@ -pub mod blobs; pub mod query; pub mod schema; -pub mod summary; + +pub use query::insert_message; + +pub fn open_or_create(session_dir: &std::path::Path) -> anyhow::Result { + let path = session_dir.join("messages.sqlite"); + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent)?; + } + let conn = rusqlite::Connection::open(&path)?; + schema::init_schema(&conn)?; + Ok(conn) +} diff --git a/src/model/msglog/query.rs b/src/model/msglog/query.rs index 1fb8c32..236bd4d 100644 --- a/src/model/msglog/query.rs +++ b/src/model/msglog/query.rs @@ -22,47 +22,3 @@ pub fn insert_message(conn: &Connection, session_id: &str, msg: &ChatMessage) -> )?; Ok(conn.last_insert_rowid()) } - -pub fn query_messages(conn: &Connection, session_id: &str, limit: usize, offset: usize) -> Result> { - let mut stmt = conn.prepare( - "SELECT role, content, tool_call_id, tool_name, tool_arguments FROM messages WHERE session_id = ?1 ORDER BY id ASC LIMIT ?2 OFFSET ?3" - )?; - let rows = stmt.query_map(params![session_id, limit as i64, offset as i64], |row| { - let role_str: String = row.get(0)?; - let content: Option = row.get(1)?; - let tool_call_id: Option = row.get(2)?; - let tool_name: Option = row.get(3)?; - let tool_arguments: Option = row.get(4)?; - let role = match role_str.as_str() { - "user" => Role::User, - "assistant" => Role::Assistant, - "system" => Role::System, - "tool" => Role::Tool, - _ => Role::User, - }; - let tool_calls = tool_arguments.and_then(|args| { - serde_json::from_str(&args).ok() - }); - Ok(ChatMessage { - role, - content, - tool_calls, - tool_call_id, - name: tool_name, - }) - })?; - let mut messages = Vec::new(); - for row in rows { - messages.push(row?); - } - Ok(messages) -} - -pub fn count_messages(conn: &Connection, session_id: &str) -> Result { - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM messages WHERE session_id = ?1", - params![session_id], - |row| row.get(0), - )?; - Ok(count) -} diff --git a/src/model/session.rs b/src/model/session.rs index 94ac9c2..2b3cfb8 100644 --- a/src/model/session.rs +++ b/src/model/session.rs @@ -41,14 +41,6 @@ impl Session { self.session_dir(base_dir).join("conversation.json") } - pub fn edit_log_path(&self, base_dir: &Path) -> PathBuf { - self.session_dir(base_dir).join("edits.jsonl") - } - - pub fn msglog_path(&self, base_dir: &Path) -> PathBuf { - self.session_dir(base_dir).join("msglog.db") - } - pub fn save(&self, base_dir: &Path) -> std::io::Result<()> { let dir = self.session_dir(base_dir); std::fs::create_dir_all(&dir)?; diff --git a/src/resources.rs b/src/resources.rs index 8aef568..4dbc6d5 100644 --- a/src/resources.rs +++ b/src/resources.rs @@ -1,16 +1,6 @@ pub const SYSTEM_PROMPT: &str = include_str!("../src-misc/system-prompt.txt"); pub const SYSTEM_TOOLS: &str = include_str!("../src-misc/system-tools.txt"); -pub const BANNER: &str = r" -███████╗███████╗███████╗██████╗ ███████╗██╗ ██╗ -╚══███╔╝██╔════╝██╔════╝██╔══██╗██╔════╝╚██╗██╔╝ - ███╔╝ █████╗ ███████╗██║ ██║█████╗ ╚███╔╝ - ███╔╝ ██╔══╝ ╚════██║██║ ██║██╔══╝ ██╔██╗ -███████╗███████╗███████║██████╔╝███████╗██╔╝ ██╗ -╚══════╝╚══════╝╚══════╝╚═════╝ ╚══════╝╚═╝ ╚═╝ - Autonomous Agentic Shell -"; - pub const HELP_TEXT: &str = " ZESDEX - Help ============= @@ -46,4 +36,4 @@ Commands: Any text is sent to the AI assistant as a prompt. File paths use workspace-relative notation. Use [0]/path for multi-workspace setups. -"; + "; diff --git a/src/service/oauth/loopback.rs b/src/service/oauth/loopback.rs index 7da7150..b420e15 100644 --- a/src/service/oauth/loopback.rs +++ b/src/service/oauth/loopback.rs @@ -13,10 +13,6 @@ impl LoopbackServer { Ok(LoopbackServer { listener, port }) } - pub fn port(&self) -> u16 { - self.port - } - pub fn redirect_uri(&self) -> String { format!("http://127.0.0.1:{}/callback", self.port) } diff --git a/src/service/oauth/manager.rs b/src/service/oauth/manager.rs index 700c059..561afa7 100644 --- a/src/service/oauth/manager.rs +++ b/src/service/oauth/manager.rs @@ -10,15 +10,6 @@ pub struct OAuthToken { } impl OAuthToken { - pub fn is_expired(&self) -> bool { - let now = SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_secs(); - now >= self.expires_at - } - - pub fn remaining_secs(&self) -> i64 { - let now = SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_secs(); - self.expires_at as i64 - now as i64 - } } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -92,52 +83,6 @@ impl OAuthManager { Ok(()) } - pub fn refresh_token(&mut self) -> Result<(), String> { - let refresh_token = self.token.as_ref() - .and_then(|t| t.refresh_token.clone()) - .ok_or("no refresh token available")?; - - let mut params = std::collections::HashMap::new(); - params.insert("grant_type", "refresh_token"); - params.insert("refresh_token", &refresh_token); - params.insert("client_id", &self.config.client_id); - - let resp = self.client - .post(&self.config.token_url) - .form(¶ms) - .send() - .map_err(|e| format!("refresh failed: {}", e))?; - - let status = resp.status(); - let body: serde_json::Value = resp.json().map_err(|e| format!("parse failed: {}", e))?; - - if !status.is_success() { - return Err(format!("refresh endpoint returned {}: {}", status, body)); - } - - let access_token = body["access_token"].as_str().ok_or("missing access_token")?.to_string(); - let expires_in = body["expires_in"].as_u64().unwrap_or(3600); - let now = SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_secs(); - - self.token = Some(OAuthToken { - access_token, - refresh_token: body["refresh_token"].as_str().map(|s| s.to_string()).or(self.token.as_ref().and_then(|t| t.refresh_token.clone())), - expires_at: now + expires_in, - token_type: body["token_type"].as_str().unwrap_or("Bearer").to_string(), - }); - - Ok(()) - } - - pub fn ensure_token(&mut self) -> Result<(), String> { - if let Some(ref token) = self.token { - if token.remaining_secs() < 60 { - return self.refresh_token(); - } - } - Ok(()) - } - pub fn build_auth_url(&self, redirect_uri: &str, state: &str, code_challenge: &str) -> String { let mut url = url::Url::parse(&self.config.auth_url).unwrap_or_else(|_| url::Url::parse("https://example.com").unwrap()); url.query_pairs_mut() diff --git a/src/tool/internet/search.rs b/src/tool/internet/search.rs index c296238..4a0ebbd 100644 --- a/src/tool/internet/search.rs +++ b/src/tool/internet/search.rs @@ -1,8 +1,37 @@ +use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; use anyhow::{Result, anyhow}; use super::super::Tool; use super::super::ToolCtx; +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum SearchProvider { + Tavily, + Brave, + SerpApi, + Google, +} + +impl SearchProvider { + pub fn from_str(s: &str) -> Option { + match s.to_lowercase().as_str() { + "tavily" => Some(SearchProvider::Tavily), + "brave" => Some(SearchProvider::Brave), + "serpapi" | "serp_api" => Some(SearchProvider::SerpApi), + "google" => Some(SearchProvider::Google), + _ => None, + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SearchResult { + pub title: String, + pub url: String, + pub snippet: String, +} + pub struct Search; impl Tool for Search { @@ -11,7 +40,7 @@ impl Tool for Search { } fn description(&self) -> &'static str { - "Search the web for information. Uses configured search provider." + "Search the web for information using a configured search provider (Tavily, Brave, SerpAPI, or Google)." } fn parameters(&self) -> Value { @@ -21,6 +50,11 @@ impl Tool for Search { "query": { "type": "string", "description": "Search query" + }, + "num_results": { + "type": "integer", + "description": "Number of results to return (default: 5)", + "default": 5 } }, "required": ["query"] @@ -35,16 +69,174 @@ impl Tool for Search { .and_then(|v| v.as_str()) .ok_or_else(|| anyhow!("missing required argument: query"))? .to_string(); - let results = mock_search(&query); - Ok(results) + let num_results = args.get("num_results") + .and_then(|v| v.as_u64()) + .unwrap_or(5) as usize; + + let provider = detect_search_provider(); + match provider { + Some(p) => search_with_provider(&p, &query, num_results), + None => Ok(format!( + "No search provider configured for query '{}'.\n\ + Set ZESDEX_SEARCH_PROVIDER and corresponding API key env vars.\n\ + Supported: tavily (ZESDEX_TAVILY_API_KEY), \ + brave (ZESDEX_BRAVE_API_KEY), \ + serpapi (ZESDEX_SERPAPI_KEY), \ + google (ZESDEX_GOOGLE_API_KEY).", + query + )), + } } } -pub(crate) fn mock_search(query: &str) -> String { - format!( - "Search results for '{}':\n\n\ - No search provider configured. Results are unavailable.\n\ - To enable web search, configure a search provider in settings.\n\ - Supported providers: tavily, brave, serpapi, google.", query - ) +fn detect_search_provider() -> Option { + if std::env::var("ZESDEX_SEARCH_PROVIDER").ok().is_some() { + let provider_str = std::env::var("ZESDEX_SEARCH_PROVIDER").unwrap_or_default(); + if let Some(p) = SearchProvider::from_str(&provider_str) { + return Some(p); + } + } + if std::env::var("ZESDEX_TAVILY_API_KEY").ok().filter(|k| !k.is_empty()).is_some() { + return Some(SearchProvider::Tavily); + } + if std::env::var("ZESDEX_BRAVE_API_KEY").ok().filter(|k| !k.is_empty()).is_some() { + return Some(SearchProvider::Brave); + } + if std::env::var("ZESDEX_SERPAPI_KEY").ok().filter(|k| !k.is_empty()).is_some() { + return Some(SearchProvider::SerpApi); + } + if std::env::var("ZESDEX_GOOGLE_API_KEY").ok().filter(|k| !k.is_empty()).is_some() { + return Some(SearchProvider::Google); + } + None +} + +fn search_with_provider(provider: &SearchProvider, query: &str, num_results: usize) -> Result { + let results = match provider { + SearchProvider::Tavily => search_tavily(query, num_results)?, + SearchProvider::Brave => search_brave(query, num_results)?, + SearchProvider::SerpApi => search_serpapi(query, num_results)?, + SearchProvider::Google => search_google(query, num_results)?, + }; + if results.is_empty() { + return Ok(format!("No results found for '{}'.", query)); + } + let mut output = format!("Search results for '{}':\n\n", query); + for (i, r) in results.iter().enumerate() { + output.push_str(&format!("{}. {}\n {}\n {}\n\n", i + 1, r.title, r.url, r.snippet)); + } + Ok(output) +} + +fn search_tavily(query: &str, num_results: usize) -> Result> { + let api_key = std::env::var("ZESDEX_TAVILY_API_KEY") + .map_err(|_| anyhow!("ZESDEX_TAVILY_API_KEY not set"))?; + let client = reqwest::blocking::Client::builder() + .timeout(std::time::Duration::from_secs(15)) + .build()?; + let body = json!({ + "api_key": api_key, + "query": query, + "max_results": num_results, + "include_answer": false, + "search_depth": "basic", + }); + let resp = client.post("https://api.tavily.com/search") + .header("Content-Type", "application/json") + .json(&body) + .send()?; + if !resp.status().is_success() { + anyhow::bail!("Tavily API error: {}", resp.status()); + } + let data: Value = resp.json()?; + let results = data["results"].as_array().cloned().unwrap_or_default(); + Ok(results.iter().filter_map(|r| { + Some(SearchResult { + title: r["title"].as_str()?.to_string(), + url: r["url"].as_str()?.to_string(), + snippet: r["content"].as_str().unwrap_or("").to_string(), + }) + }).collect()) +} + +fn search_brave(query: &str, num_results: usize) -> Result> { + let api_key = std::env::var("ZESDEX_BRAVE_API_KEY") + .map_err(|_| anyhow!("ZESDEX_BRAVE_API_KEY not set"))?; + let client = reqwest::blocking::Client::builder() + .timeout(std::time::Duration::from_secs(15)) + .build()?; + let resp = client.get("https://api.search.brave.com/res/v1/web/search") + .header("Accept", "application/json") + .header("Accept-Encoding", "gzip") + .header("X-Subscription-Token", &api_key) + .query(&[("q", query), ("count", &num_results.to_string())]) + .send()?; + if !resp.status().is_success() { + anyhow::bail!("Brave API error: {}", resp.status()); + } + let data: Value = resp.json()?; + let results = data["web"]["results"].as_array().cloned().unwrap_or_default(); + Ok(results.iter().filter_map(|r| { + Some(SearchResult { + title: r["title"].as_str()?.to_string(), + url: r["url"].as_str()?.to_string(), + snippet: r["description"].as_str().unwrap_or("").to_string(), + }) + }).collect()) +} + +fn search_serpapi(query: &str, num_results: usize) -> Result> { + let api_key = std::env::var("ZESDEX_SERPAPI_KEY") + .map_err(|_| anyhow!("ZESDEX_SERPAPI_KEY not set"))?; + let client = reqwest::blocking::Client::builder() + .timeout(std::time::Duration::from_secs(15)) + .build()?; + let resp = client.get("https://serpapi.com/search.json") + .query(&[ + ("q", query), + ("api_key", &api_key), + ("engine", "google"), + ("num", &num_results.to_string()), + ]) + .send()?; + if !resp.status().is_success() { + anyhow::bail!("SerpAPI error: {}", resp.status()); + } + let data: Value = resp.json()?; + let results = data["organic_results"].as_array().cloned().unwrap_or_default(); + Ok(results.iter().filter_map(|r| { + let title = r.get("title")?.as_str()?.to_string(); + let url = r.get("link")?.as_str()?.to_string(); + let snippet = r.get("snippet").and_then(|s| s.as_str()).unwrap_or("").to_string(); + Some(SearchResult { title, url, snippet }) + }).collect()) +} + +fn search_google(query: &str, num_results: usize) -> Result> { + let api_key = std::env::var("ZESDEX_GOOGLE_API_KEY") + .map_err(|_| anyhow!("ZESDEX_GOOGLE_API_KEY not set"))?; + let cx = std::env::var("ZESDEX_GOOGLE_CX") + .map_err(|_| anyhow!("ZESDEX_GOOGLE_CX (Custom Search Engine ID) not set"))?; + let client = reqwest::blocking::Client::builder() + .timeout(std::time::Duration::from_secs(15)) + .build()?; + let resp = client.get("https://www.googleapis.com/customsearch/v1") + .query(&[ + ("q", query), + ("key", &api_key), + ("cx", &cx), + ("num", &num_results.min(10).to_string()), + ]) + .send()?; + if !resp.status().is_success() { + anyhow::bail!("Google Custom Search API error: {}", resp.status()); + } + let data: Value = resp.json()?; + let results = data["items"].as_array().cloned().unwrap_or_default(); + Ok(results.iter().filter_map(|r| { + let title = r.get("title")?.as_str()?.to_string(); + let url = r.get("link")?.as_str()?.to_string(); + let snippet = r.get("snippet").and_then(|s| s.as_str()).unwrap_or("").to_string(); + Some(SearchResult { title, url, snippet }) + }).collect()) } diff --git a/src/tool/mod.rs b/src/tool/mod.rs index 47d6cae..f8fa3a4 100644 --- a/src/tool/mod.rs +++ b/src/tool/mod.rs @@ -13,7 +13,6 @@ pub mod plan; pub mod search; pub mod seqthink; pub mod shell; -pub mod shell_filter; pub mod utility; pub mod workflow; @@ -36,7 +35,7 @@ pub struct ToolCtx { pub workspaces: Vec, pub session_dir: PathBuf, pub memory_dir: PathBuf, - pub download_dir: PathBuf, + pub _download_dir: PathBuf, pub worktrees_dir: PathBuf, pub dir_cache: std::sync::Arc>, pub internet_mode: super::model::settings::InternetMode, @@ -89,20 +88,14 @@ impl Default for ToolCtxBuilder { } impl ToolCtxBuilder { - pub fn workspaces(mut self, v: Vec) -> Self { self.workspaces = v; self } pub fn session_dir(mut self, v: PathBuf) -> Self { self.session_dir = v; self } - pub fn memory_dir(mut self, v: PathBuf) -> Self { self.memory_dir = v; self } - pub fn download_dir(mut self, v: PathBuf) -> Self { self.download_dir = v; self } - pub fn worktrees_dir(mut self, v: PathBuf) -> Self { self.worktrees_dir = v; self } - pub fn internet_mode(mut self, v: super::model::settings::InternetMode) -> Self { self.internet_mode = v; self } pub fn origin(mut self, v: crate::app::state::types::Origin) -> Self { self.origin = v; self } - pub fn graduated_checks(mut self, v: Vec) -> Self { self.graduated_checks = v; self } pub fn build(self) -> ToolCtx { ToolCtx { workspaces: self.workspaces, session_dir: self.session_dir, memory_dir: self.memory_dir, - download_dir: self.download_dir, + _download_dir: self.download_dir, worktrees_dir: self.worktrees_dir, dir_cache: self.dir_cache, internet_mode: self.internet_mode, @@ -130,6 +123,7 @@ pub fn all_tools() -> Vec> { Box::new(super::tool::plan::PlanEnter), Box::new(super::tool::plan::PlanReady), Box::new(super::tool::workflow::WorkflowRun), + Box::new(super::tool::workflow::NoteFinding), Box::new(super::tool::internet::fetch::Fetch), Box::new(super::tool::internet::download::Download), Box::new(super::tool::internet::search::Search), @@ -162,11 +156,6 @@ pub fn tool_defs(tools: &[Box]) -> Vec Result { let _parts: Vec<&str> = rel.splitn(2, '/').collect(); let (ws_idx, path) = if rel.starts_with('[') { diff --git a/src/tool/workflow.rs b/src/tool/workflow.rs index 8f8b594..c6b6180 100644 --- a/src/tool/workflow.rs +++ b/src/tool/workflow.rs @@ -11,7 +11,7 @@ impl Tool for WorkflowRun { } fn description(&self) -> &'static str { - "Execute a workflow script by delegating to the workflow engine" + "Execute a workflow script that can spawn multiple subagents in parallel, pipeline, or phased stages. Use when a task benefits from decomposition into independent subtasks. Simple tasks should be handled inline without this tool." } fn parameters(&self) -> Value { @@ -20,11 +20,11 @@ impl Tool for WorkflowRun { "properties": { "script": { "type": "string", - "description": "Workflow script content or path to a workflow file" + "description": "JSON-encoded workflow script with name, description, script (Agent/Parallel/Pipeline/Phase primitives), and options (max_concurrency, continue_on_error)" }, "args": { "type": "object", - "description": "Optional arguments passed to the workflow script" + "description": "Optional string key-value arguments passed to the workflow script for template substitution ({{key}} placeholders)" } }, "required": ["script"] @@ -52,3 +52,37 @@ impl Tool for WorkflowRun { crate::app::workflow::engine::run_workflow(&workflow_script, &workflow_args) } } + +pub struct NoteFinding; + +impl Tool for NoteFinding { + fn name(&self) -> &'static str { + "note_finding" + } + + fn description(&self) -> &'static str { + "Share a finding with sibling agents in the same workflow_run. Findings are ephemeral to the current run and will be prepended to other agents' next tool-round context. Does not persist to memory." + } + + fn parameters(&self) -> Value { + json!({ + "type": "object", + "properties": { + "text": { + "type": "string", + "description": "The finding to share with sibling agents" + } + }, + "required": ["text"] + }) + } + + fn run(&self, _ctx: &ToolCtx, args: &Value) -> Result { + let text = args.get("text") + .and_then(|v| v.as_str()) + .ok_or_else(|| anyhow!("missing required argument: text"))?; + + crate::app::workflow::engine::note_finding(text); + Ok(format!("finding recorded: {}", text.chars().take(80).collect::())) + } +} diff --git a/src/view/chat.rs b/src/view/chat.rs index 708219d..903bd30 100644 --- a/src/view/chat.rs +++ b/src/view/chat.rs @@ -38,23 +38,9 @@ pub fn draw_chat(frame: &mut Frame, area: Rect, state: &crate::app::state::rest: let mut display_lines: Vec = Vec::new(); - if scroll_offset > 0 { - let above = messages.len().saturating_sub(scroll_offset + max_visible).saturating_sub(1); - if above > 0 { - display_lines.push(Line::from(Span::styled( - format!(" ↑ {} more messages above", above), - Style::default().fg(Theme::DIM), - ))); - } - } + let _total_msgs = messages.len(); - let iter_start = if scroll_offset + max_visible >= messages.len() { - 0usize - } else { - messages.len().saturating_sub(scroll_offset + max_visible) - }; - - for msg in messages.iter().skip(iter_start).take(scroll_offset + max_visible) { + for msg in messages.iter() { let role_color = match msg.role { crate::dto::chat::message::Role::User => Theme::ROLE_USER, crate::dto::chat::message::Role::Assistant => Theme::ROLE_ASSISTANT, @@ -98,23 +84,12 @@ pub fn draw_chat(frame: &mut Frame, area: Rect, state: &crate::app::state::rest: display_lines.push(Line::from(Span::raw(""))); } - if state.turn_in_flight() { - let last_is_user = messages.last().map(|m| matches!(m.role, crate::dto::chat::message::Role::User)).unwrap_or(false); - if last_is_user { - display_lines.push(Line::from(vec![ - Span::styled(" AI ", Style::default().fg(Theme::BG).bg(Theme::ROLE_ASSISTANT).add_modifier(Modifier::BOLD)), - Span::styled(" Thinking...", Style::default().fg(Theme::DIM)), - ])); - display_lines.push(Line::from(Span::raw(""))); - } - } - - if messages.len() > scroll_offset + max_visible { - let below = messages.len().saturating_sub(scroll_offset + max_visible); - display_lines.push(Line::from(Span::styled( - format!(" ↓ {} more messages below", below), - Style::default().fg(Theme::DIM), - ))); + if state.misc.thinking { + display_lines.push(Line::from(vec![ + Span::styled(" AI ", Style::default().fg(Theme::BG).bg(Theme::ROLE_ASSISTANT).add_modifier(Modifier::BOLD)), + Span::styled(" Thinking...", Style::default().fg(Theme::DIM)), + ])); + display_lines.push(Line::from(Span::raw(""))); } let mut title = String::from(" Chat "); @@ -128,13 +103,15 @@ pub fn draw_chat(frame: &mut Frame, area: Rect, state: &crate::app::state::rest: .title(title); let total = display_lines.len(); - let end_idx = total.saturating_sub(scroll_offset); + let max_offset = total.saturating_sub(max_visible); + let offset = scroll_offset.min(max_offset); + + let end_idx = total.saturating_sub(offset); let start_idx = end_idx.saturating_sub(max_visible); - let end_idx = end_idx.min(total); - let visible: Vec = if start_idx < end_idx { + let visible: Vec = if start_idx < end_idx && start_idx < total { display_lines[start_idx..end_idx].to_vec() } else { - Vec::new() + display_lines[total.saturating_sub(max_visible)..total].to_vec() }; let paragraph = Paragraph::new(visible) diff --git a/src/view/mod.rs b/src/view/mod.rs index fb578f7..a009e41 100644 --- a/src/view/mod.rs +++ b/src/view/mod.rs @@ -1,5 +1,4 @@ pub mod chat; -pub mod markdown; pub mod status; pub mod theme; pub mod workflow; diff --git a/src/view/theme.rs b/src/view/theme.rs index 45e7694..ebf4bd4 100644 --- a/src/view/theme.rs +++ b/src/view/theme.rs @@ -4,7 +4,6 @@ pub struct Theme; impl Theme { pub const PRIMARY: Color = Color::Cyan; - pub const SECONDARY: Color = Color::Magenta; pub const SUCCESS: Color = Color::Green; pub const WARNING: Color = Color::Yellow; pub const ERROR: Color = Color::Red;