agentspec_provider/local/
claude.rs1use 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
12pub(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
271pub(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}