Skip to main content

agentspec_provider/local/
claude.rs

1use crate::error::ProviderError;
2use crate::ir::{Capability, GIT_READONLY_COMMANDS, ProviderConfig, Sandbox, SandboxTranslator};
3use crate::types::{AiEvent, AiProvider, AiRequest, AiResponse, AiUsage};
4use anyhow::{Context, Result};
5use async_trait::async_trait;
6use tokio::io::{AsyncBufReadExt, BufReader};
7use tokio::process::Command;
8use tokio::sync::mpsc::UnboundedSender;
9
10use super::json::embed_schema;
11
12/// Configuration for the Claude CLI provider.
13pub(crate) struct ClaudeConfig {
14    pub model: Option<String>,
15    pub budget: f64,
16    pub sandbox: Option<Sandbox>,
17    pub debug: bool,
18}
19
20impl Default for ClaudeConfig {
21    fn default() -> Self {
22        Self {
23            model: None,
24            budget: 0.50,
25            sandbox: None,
26            debug: false,
27        }
28    }
29}
30
31pub struct ClaudeProvider {
32    config: ClaudeConfig,
33}
34
35impl ClaudeProvider {
36    pub(crate) fn new(config: ClaudeConfig) -> Self {
37        Self { config }
38    }
39
40    pub(crate) fn from_provider_config(config: &ProviderConfig) -> Self {
41        Self::new(ClaudeConfig {
42            model: config.model.clone(),
43            budget: config.budget.unwrap_or(0.50),
44            sandbox: config.sandbox.clone(),
45            debug: config.debug,
46        })
47    }
48
49    fn base_command(&self, working_dir: &str) -> Command {
50        let model = self.config.model.as_deref().unwrap_or("haiku");
51        let mut cmd = Command::new("claude");
52        cmd.current_dir(working_dir).arg("--model").arg(model);
53
54        if let Some(sandbox) = &self.config.sandbox {
55            for tool in self.translate_sandbox(sandbox) {
56                cmd.arg("--allowed-tools").arg(tool);
57            }
58        }
59
60        cmd.arg("--max-budget-usd")
61            .arg(format!("{:.2}", self.config.budget))
62            .arg("-p");
63        cmd
64    }
65
66    async fn request_streaming(
67        &self,
68        req: &AiRequest,
69        events: UnboundedSender<AiEvent>,
70    ) -> Result<AiResponse> {
71        let system = embed_schema(&req.system_prompt, req.json_schema.as_deref());
72
73        let mut cmd = self.base_command(&req.working_dir);
74        cmd.arg(&req.user_prompt)
75            .arg("--system-prompt")
76            .arg(&system)
77            .arg("--output-format")
78            .arg("stream-json")
79            .arg("--verbose");
80
81        if self.config.debug {
82            eprintln!(
83                "[DEBUG] claude stream-json (model={}, budget={:.2})",
84                self.config.model.as_deref().unwrap_or("haiku"),
85                self.config.budget
86            );
87        }
88
89        let mut child = cmd
90            .stdout(std::process::Stdio::piped())
91            .stderr(std::process::Stdio::piped())
92            .spawn()
93            .context("failed to run claude CLI")?;
94
95        let stdout = child.stdout.take().unwrap();
96        let stderr_handle = child.stderr.take().unwrap();
97
98        let stderr_task = tokio::spawn(async move {
99            let mut buf = String::new();
100            let _ = tokio::io::AsyncReadExt::read_to_string(
101                &mut BufReader::new(stderr_handle),
102                &mut buf,
103            )
104            .await;
105            buf
106        });
107
108        let mut reader = BufReader::new(stdout).lines();
109        let mut result_text = String::new();
110        let mut usage = None;
111
112        while let Ok(Some(line)) = reader.next_line().await {
113            let event: serde_json::Value = match serde_json::from_str(line.trim()) {
114                Ok(v) => v,
115                Err(_) => continue,
116            };
117
118            parse_tool_calls(&event, &events);
119
120            if event.get("type").and_then(|t| t.as_str()) == Some("result") {
121                if let Some(r) = event.get("result") {
122                    let raw = match r {
123                        serde_json::Value::String(s) => s.clone(),
124                        _ => r.to_string(),
125                    };
126                    result_text = super::json::extract_json(&raw).unwrap_or(raw);
127                }
128                usage = extract_usage(&event);
129            }
130        }
131
132        let stderr_text = stderr_task.await.unwrap_or_default();
133        let status = child.wait().await?;
134
135        if !status.success() {
136            anyhow::bail!(ProviderError::BackendFailed(format!(
137                "claude CLI failed (exit {}): {}",
138                status,
139                stderr_text.trim()
140            )));
141        }
142
143        if result_text.is_empty() {
144            anyhow::bail!(ProviderError::ParseResponse(
145                "no result in claude stream".into()
146            ));
147        }
148
149        Ok(AiResponse {
150            text: result_text,
151            usage,
152        })
153    }
154
155    async fn request_batch(&self, req: &AiRequest) -> Result<AiResponse> {
156        let mut cmd = self.base_command(&req.working_dir);
157        cmd.arg(&req.user_prompt)
158            .arg("--system-prompt")
159            .arg(&req.system_prompt)
160            .arg("--output-format")
161            .arg("json");
162
163        if let Some(schema) = &req.json_schema {
164            cmd.arg("--json-schema").arg(schema);
165        }
166
167        if self.config.debug {
168            eprintln!(
169                "[DEBUG] claude json (model={}, budget={:.2})",
170                self.config.model.as_deref().unwrap_or("haiku"),
171                self.config.budget
172            );
173        }
174
175        let output = cmd.output().await.context("failed to run claude CLI")?;
176        let raw = String::from_utf8_lossy(&output.stdout).to_string();
177        let stderr = String::from_utf8_lossy(&output.stderr);
178
179        if self.config.debug {
180            eprintln!("[DEBUG] exit: {}", output.status);
181            eprintln!("[DEBUG] stdout (first 500): {}", &raw[..raw.len().min(500)]);
182            if !stderr.is_empty() {
183                eprintln!("[DEBUG] stderr: {stderr}");
184            }
185        }
186
187        if !output.status.success() {
188            anyhow::bail!(ProviderError::BackendFailed(format!(
189                "claude CLI failed (exit {}): {}",
190                output.status,
191                stderr.trim()
192            )));
193        }
194
195        let parsed: serde_json::Value =
196            serde_json::from_str(&raw).context("failed to parse claude JSON response")?;
197        let usage = extract_usage(&parsed);
198
199        if req.json_schema.is_some() {
200            let structured = &parsed["structured_output"];
201            if structured.is_null() {
202                anyhow::bail!(ProviderError::ParseResponse(
203                    "empty structured_output from claude".into()
204                ));
205            }
206            Ok(AiResponse {
207                text: structured.to_string(),
208                usage,
209            })
210        } else {
211            let text = parsed
212                .get("result")
213                .map(|r| match r {
214                    serde_json::Value::String(s) => s.clone(),
215                    _ => r.to_string(),
216                })
217                .unwrap_or(raw);
218            Ok(AiResponse { text, usage })
219        }
220    }
221}
222
223impl SandboxTranslator for ClaudeProvider {
224    fn translate_allowed(&self, capabilities: &[Capability]) -> Vec<String> {
225        let mut tools = Vec::new();
226        for cap in capabilities {
227            match cap {
228                Capability::ReadFile => tools.push("Read".into()),
229                Capability::GitReadOnly => {
230                    for cmd in GIT_READONLY_COMMANDS {
231                        tools.push(format!("Bash(git:{cmd})"));
232                    }
233                }
234                Capability::ShellCommand { pattern } => {
235                    tools.push(format!("Bash({pattern})"));
236                }
237                Capability::Custom(s) => tools.push(s.clone()),
238                Capability::WriteFile | Capability::Network => {}
239            }
240        }
241        tools
242    }
243}
244
245#[async_trait]
246impl AiProvider for ClaudeProvider {
247    fn name(&self) -> &str {
248        "claude"
249    }
250
251    async fn is_available(&self) -> bool {
252        Command::new("claude")
253            .arg("--version")
254            .output()
255            .await
256            .is_ok_and(|o| o.status.success())
257    }
258
259    async fn request(
260        &self,
261        req: &AiRequest,
262        events: Option<tokio::sync::mpsc::UnboundedSender<AiEvent>>,
263    ) -> Result<AiResponse> {
264        match events {
265            Some(tx) => self.request_streaming(req, tx).await,
266            None => self.request_batch(req).await,
267        }
268    }
269}
270
271/// Parse tool calls from Claude's NDJSON stream events.
272pub(crate) fn parse_tool_calls(event: &serde_json::Value, events: &UnboundedSender<AiEvent>) {
273    if let Some(content) = event.pointer("/message/content")
274        && let Some(arr) = content.as_array()
275    {
276        for item in arr {
277            if item["type"] == "tool_use"
278                && let Some(input) = extract_tool_input(item)
279            {
280                let tool = item["name"].as_str().unwrap_or("unknown").to_string();
281                let _ = events.send(AiEvent::ToolCall { tool, input });
282            }
283        }
284    }
285
286    if event.get("type").and_then(|t| t.as_str()) == Some("stream_event")
287        && let Some(inner) = event.get("event")
288        && inner.get("type").and_then(|t| t.as_str()) == Some("content_block_start")
289        && let Some(block) = inner.get("content_block")
290        && block.get("type").and_then(|t| t.as_str()) == Some("tool_use")
291    {
292        let tool = block["name"].as_str().unwrap_or("unknown").to_string();
293        let input = extract_tool_input(block).unwrap_or_default();
294        if !input.is_empty() {
295            let _ = events.send(AiEvent::ToolCall { tool, input });
296        }
297    }
298}
299
300fn extract_tool_input(item: &serde_json::Value) -> Option<String> {
301    if let Some(cmd) = item.pointer("/input/command").and_then(|c| c.as_str()) {
302        return Some(cmd.to_string());
303    }
304    if let Some(path) = item.pointer("/input/file_path").and_then(|p| p.as_str()) {
305        return Some(path.to_string());
306    }
307    item.get("input")
308        .filter(|i| !i.is_null())
309        .map(|i| serde_json::to_string(i).unwrap_or_default())
310        .filter(|s| !s.is_empty() && s != "{}")
311}
312
313fn extract_usage(parsed: &serde_json::Value) -> Option<AiUsage> {
314    let u = parsed.get("usage")?;
315    Some(AiUsage {
316        input_tokens: u.get("input_tokens")?.as_u64()?,
317        output_tokens: u.get("output_tokens")?.as_u64()?,
318        cost_usd: parsed.get("cost_usd").and_then(|c| c.as_f64()),
319    })
320}
321
322#[cfg(test)]
323mod tests {
324    use super::*;
325
326    #[test]
327    fn translate_git_readonly() {
328        let provider = ClaudeProvider::new(ClaudeConfig::default());
329        let caps = vec![Capability::GitReadOnly];
330        let tools = provider.translate_allowed(&caps);
331        assert!(tools.contains(&"Bash(git:diff)".to_string()));
332        assert!(tools.contains(&"Bash(git:log)".to_string()));
333        assert!(tools.contains(&"Bash(git:blame)".to_string()));
334        assert_eq!(tools.len(), GIT_READONLY_COMMANDS.len());
335    }
336
337    #[test]
338    fn translate_read_file() {
339        let provider = ClaudeProvider::new(ClaudeConfig::default());
340        let caps = vec![Capability::ReadFile];
341        let tools = provider.translate_allowed(&caps);
342        assert_eq!(tools, vec!["Read"]);
343    }
344
345    #[test]
346    fn translate_custom_passthrough() {
347        let provider = ClaudeProvider::new(ClaudeConfig::default());
348        let caps = vec![Capability::Custom("Bash(npm:test)".into())];
349        let tools = provider.translate_allowed(&caps);
350        assert_eq!(tools, vec!["Bash(npm:test)"]);
351    }
352
353    #[test]
354    fn translate_sandbox_filters_denied() {
355        let provider = ClaudeProvider::new(ClaudeConfig::default());
356        let sandbox = Sandbox {
357            allowed: vec![Capability::GitReadOnly, Capability::ReadFile],
358            denied: vec![Capability::ReadFile],
359        };
360        let tools = provider.translate_sandbox(&sandbox);
361        assert!(!tools.contains(&"Read".to_string()));
362        assert!(tools.contains(&"Bash(git:diff)".to_string()));
363    }
364
365    #[test]
366    fn translate_empty_sandbox() {
367        let provider = ClaudeProvider::new(ClaudeConfig::default());
368        let sandbox = Sandbox::default();
369        let tools = provider.translate_sandbox(&sandbox);
370        assert!(tools.is_empty());
371    }
372}