Skip to main content

claude_code_sdk_rust/internal/
transport.rs

1use async_trait::async_trait;
2use std::collections::HashMap;
3use std::process::Stdio;
4use std::sync::Arc;
5use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
6use tokio::process::{Child, ChildStdin, ChildStdout, Command};
7use tokio::sync::Mutex;
8
9use crate::error::{CLIConnectionError, CLINotFoundError, ClaudeSDKError, ProcessError, Result};
10use crate::internal::stdout_decoder::StdoutDecoder;
11use crate::types::ClaudeAgentOptions;
12
13const DEFAULT_ENTRY_POINT: &str = "sdk-rust";
14const DEFAULT_MAX_BUFFER_SIZE: usize = 1024 * 1024;
15
16#[async_trait]
17pub trait Transport: Send + Sync {
18    async fn connect(&mut self) -> Result<()>;
19    async fn write(&mut self, data: &[u8]) -> Result<()>;
20    async fn close_input(&mut self) -> Result<()>;
21    async fn read(&mut self) -> Result<Option<Vec<u8>>>;
22    async fn close(&mut self) -> Result<()>;
23}
24
25#[derive(Debug)]
26pub struct SubprocessCLITransport {
27    options: TransportOptions,
28    child: Option<Child>,
29    stdin: Option<ChildStdin>,
30    stdout_reader: Option<BufReader<ChildStdout>>,
31    stdout_decoder: StdoutDecoder,
32    stderr: Arc<Mutex<String>>,
33}
34
35#[derive(Debug, Clone)]
36pub struct TransportOptions {
37    pub tools: Vec<String>,
38    pub tools_set: bool,
39    pub tools_preset: Option<crate::types::ToolsPreset>,
40    pub allowed_tools: Vec<String>,
41    pub system_prompt: Option<String>,
42    pub system_prompt_preset: Option<crate::types::SystemPromptPreset>,
43    pub system_prompt_file: Option<crate::types::SystemPromptFile>,
44    pub mcp_servers: std::collections::HashMap<String, crate::types::MCPServerConfig>,
45    pub mcp_servers_config: Option<String>,
46    pub permission_mode: Option<crate::types::PermissionMode>,
47    pub continue_conversation: bool,
48    pub resume: Option<String>,
49    pub session_id: Option<String>,
50    pub fork_session: bool,
51    pub max_turns: Option<i32>,
52    pub max_budget_usd: Option<f64>,
53    pub task_budget: Option<crate::types::TaskBudget>,
54    pub disallowed_tools: Vec<String>,
55    pub model: Option<String>,
56    pub fallback_model: Option<String>,
57    pub betas: Vec<crate::types::SdkBeta>,
58    pub permission_prompt_tool_name: Option<String>,
59    pub cwd: Option<String>,
60    pub cli_path: Option<String>,
61    pub settings: Option<String>,
62    pub add_dirs: Vec<String>,
63    pub env: std::collections::HashMap<String, String>,
64    pub extra_args: std::collections::HashMap<String, Option<String>>,
65    pub max_buffer_size: Option<usize>,
66    pub user: Option<String>,
67    pub include_partial_messages: bool,
68    pub include_hook_events: bool,
69    pub strict_mcp_config: bool,
70    pub setting_sources: Option<Vec<crate::types::SettingSource>>,
71    pub skills: Option<crate::types::SkillsConfig>,
72    pub sandbox: Option<crate::types::SandboxSettings>,
73    pub plugins: Vec<crate::types::SDKPluginConfig>,
74    pub max_thinking_tokens: Option<i32>,
75    pub thinking: Option<crate::types::ThinkingConfig>,
76    pub effort: Option<crate::types::EffortLevel>,
77    pub output_format: Option<serde_json::Map<String, serde_json::Value>>,
78    pub enable_file_checkpointing: bool,
79    pub stderr: Option<crate::types::StderrCallback>,
80    pub can_use_tool: Option<crate::types::CanUseToolCallback>,
81    pub sdk_mcp_servers: std::collections::HashMap<String, crate::mcp::SimpleMCPServer>,
82    pub session_store_enabled: bool,
83}
84
85impl From<&ClaudeAgentOptions> for TransportOptions {
86    fn from(opts: &ClaudeAgentOptions) -> Self {
87        Self {
88            tools: opts.tools.clone(),
89            tools_set: opts.tools_set,
90            tools_preset: opts.tools_preset.clone(),
91            allowed_tools: opts.allowed_tools.clone(),
92            system_prompt: opts.system_prompt.clone(),
93            system_prompt_preset: opts.system_prompt_preset.clone(),
94            system_prompt_file: opts.system_prompt_file.clone(),
95            mcp_servers: opts.mcp_servers.clone(),
96            mcp_servers_config: opts.mcp_servers_config.clone(),
97            permission_mode: opts.permission_mode,
98            continue_conversation: opts.continue_conversation,
99            resume: opts.resume.clone(),
100            session_id: opts.session_id.clone(),
101            fork_session: opts.fork_session,
102            max_turns: opts.max_turns,
103            max_budget_usd: opts.max_budget_usd,
104            task_budget: opts.task_budget.clone(),
105            disallowed_tools: opts.disallowed_tools.clone(),
106            model: opts.model.clone(),
107            fallback_model: opts.fallback_model.clone(),
108            betas: opts.betas.clone(),
109            permission_prompt_tool_name: opts
110                .permission_prompt_tool_name
111                .clone()
112                .or_else(|| opts.can_use_tool.as_ref().map(|_| "stdio".to_string())),
113            cwd: opts.cwd.clone(),
114            cli_path: opts.cli_path.clone(),
115            settings: opts.settings.clone(),
116            add_dirs: opts.add_dirs.clone(),
117            env: opts.env.clone(),
118            extra_args: opts.extra_args.clone(),
119            max_buffer_size: opts.max_buffer_size,
120            user: opts.user.clone(),
121            include_partial_messages: opts.include_partial_messages,
122            include_hook_events: opts.include_hook_events,
123            strict_mcp_config: opts.strict_mcp_config,
124            setting_sources: opts.setting_sources.clone(),
125            skills: opts.skills.clone(),
126            sandbox: opts.sandbox.clone(),
127            plugins: opts.plugins.clone(),
128            max_thinking_tokens: opts.max_thinking_tokens,
129            thinking: opts.thinking.clone(),
130            effort: opts.effort.clone(),
131            output_format: opts.output_format.clone(),
132            enable_file_checkpointing: opts.enable_file_checkpointing,
133            stderr: opts.stderr.clone(),
134            can_use_tool: opts.can_use_tool.clone(),
135            sdk_mcp_servers: opts.sdk_mcp_servers.clone(),
136            session_store_enabled: opts.session_store.is_some(),
137        }
138    }
139}
140
141impl SubprocessCLITransport {
142    pub fn new(options: TransportOptions) -> Self {
143        let max_buffer_size = options.max_buffer_size.unwrap_or(DEFAULT_MAX_BUFFER_SIZE);
144        Self {
145            options,
146            child: None,
147            stdin: None,
148            stdout_reader: None,
149            stdout_decoder: StdoutDecoder::new(max_buffer_size),
150            stderr: Arc::new(Mutex::new(String::new())),
151        }
152    }
153
154    fn resolve_cli_path(&self) -> Result<String> {
155        crate::internal::cli_discovery::find_cli_path(self.options.cli_path.as_deref())
156    }
157
158    fn build_args(&self) -> Result<Vec<String>> {
159        crate::internal::cli_args::build_cli_args(&self.options)
160    }
161
162    fn build_env(&self) -> std::collections::HashMap<String, String> {
163        build_process_env(std::env::vars(), &self.options)
164    }
165    async fn finish_read(&mut self) -> Result<Option<Vec<u8>>> {
166        if let Some(ref mut child) = self.child {
167            match child.wait().await {
168                Ok(status) => {
169                    if !status.success() {
170                        let stderr = self.stderr.lock().await.clone();
171                        return Err(ProcessError::new(
172                            "Claude Code process exited with error",
173                            status.code(),
174                            stderr,
175                        )
176                        .into());
177                    }
178                }
179                Err(e) => {
180                    return Err(CLIConnectionError::new(format!(
181                        "failed to wait for process: {}",
182                        e
183                    ))
184                    .into());
185                }
186            }
187        }
188        Ok(None)
189    }
190}
191
192fn build_process_env<I>(inherited: I, options: &TransportOptions) -> HashMap<String, String>
193where
194    I: IntoIterator<Item = (String, String)>,
195{
196    let mut env = inherited
197        .into_iter()
198        .filter(|(key, _)| key != "CLAUDECODE")
199        .collect::<HashMap<_, _>>();
200
201    env.insert(
202        "CLAUDE_CODE_ENTRYPOINT".to_string(),
203        DEFAULT_ENTRY_POINT.to_string(),
204    );
205
206    for (key, value) in &options.env {
207        env.insert(key.clone(), value.clone());
208    }
209
210    env.insert(
211        "CLAUDE_AGENT_SDK_VERSION".to_string(),
212        env!("CARGO_PKG_VERSION").to_string(),
213    );
214
215    apply_otel_trace_context(&mut env, &options.env, active_otel_trace_context());
216
217    if options.enable_file_checkpointing {
218        env.insert(
219            "CLAUDE_CODE_ENABLE_SDK_FILE_CHECKPOINTING".to_string(),
220            "true".to_string(),
221        );
222    }
223
224    if let Some(ref cwd) = options.cwd {
225        env.insert("PWD".to_string(), cwd.clone());
226    }
227
228    env
229}
230
231fn apply_otel_trace_context(
232    env: &mut HashMap<String, String>,
233    explicit_env: &HashMap<String, String>,
234    carrier: HashMap<String, String>,
235) {
236    if !carrier.contains_key("traceparent") {
237        return;
238    }
239
240    for key in ["TRACEPARENT", "TRACESTATE"] {
241        if !explicit_env.contains_key(key) {
242            env.remove(key);
243        }
244    }
245
246    for (key, value) in carrier {
247        let env_key = key.to_ascii_uppercase();
248        if !explicit_env.contains_key(&env_key) {
249            env.insert(env_key, value);
250        }
251    }
252}
253
254#[cfg(feature = "otel")]
255fn active_otel_trace_context() -> HashMap<String, String> {
256    let mut carrier = HashMap::new();
257    opentelemetry::global::get_text_map_propagator(|propagator| {
258        propagator.inject(&mut carrier);
259    });
260    carrier
261}
262
263#[cfg(not(feature = "otel"))]
264fn active_otel_trace_context() -> HashMap<String, String> {
265    HashMap::new()
266}
267
268#[async_trait]
269impl Transport for SubprocessCLITransport {
270    async fn connect(&mut self) -> Result<()> {
271        if self.child.is_some() {
272            return Ok(());
273        }
274
275        let cli_path = self.resolve_cli_path()?;
276        if std::env::var_os("CLAUDE_AGENT_SDK_SKIP_VERSION_CHECK").is_none() {
277            let _ = crate::internal::cli_discovery::check_cli_version(&cli_path).await;
278        }
279
280        if let Some(ref cwd) = self.options.cwd {
281            if !tokio::fs::metadata(cwd)
282                .await
283                .map(|m| m.is_dir())
284                .unwrap_or(false)
285            {
286                return Err(CLIConnectionError::new(format!(
287                    "working directory does not exist: {}",
288                    cwd
289                ))
290                .into());
291            }
292        }
293
294        let args = self.build_args()?;
295        let env = self.build_env();
296
297        let mut cmd = Command::new(&cli_path);
298        cmd.args(&args)
299            .stdin(Stdio::piped())
300            .stdout(Stdio::piped())
301            .stderr(Stdio::piped())
302            .kill_on_drop(true);
303
304        if let Some(ref cwd) = self.options.cwd {
305            cmd.current_dir(cwd);
306        }
307
308        for (key, value) in &env {
309            cmd.env(key, value);
310        }
311
312        let mut child = cmd.spawn().map_err(|e| {
313            if e.kind() == std::io::ErrorKind::NotFound {
314                ClaudeSDKError::CLINotFound(CLINotFoundError::new(
315                    "Claude Code not found",
316                    cli_path,
317                ))
318            } else {
319                CLIConnectionError::new(format!("failed to start Claude Code: {}", e)).into()
320            }
321        })?;
322
323        let stdin = child
324            .stdin
325            .take()
326            .ok_or_else(|| CLIConnectionError::new("failed to open CLI stdin"))?;
327        let stdout = child
328            .stdout
329            .take()
330            .ok_or_else(|| CLIConnectionError::new("failed to open CLI stdout"))?;
331        let stderr = child
332            .stderr
333            .take()
334            .ok_or_else(|| CLIConnectionError::new("failed to open CLI stderr"))?;
335
336        let stderr_arc = self.stderr.clone();
337        let stderr_callback = self.options.stderr.clone();
338        tokio::spawn(async move {
339            let mut reader = BufReader::new(stderr);
340            let mut line = String::new();
341            while let Ok(n) = reader.read_line(&mut line).await {
342                if n == 0 {
343                    break;
344                }
345                let mut stderr_guard = stderr_arc.lock().await;
346                stderr_guard.push_str(&line);
347                if let Some(callback) = &stderr_callback {
348                    callback.call(line.clone());
349                }
350                line.clear();
351            }
352        });
353
354        self.child = Some(child);
355        self.stdin = Some(stdin);
356        self.stdout_reader = Some(BufReader::new(stdout));
357
358        Ok(())
359    }
360
361    async fn write(&mut self, data: &[u8]) -> Result<()> {
362        let stdin = self
363            .stdin
364            .as_mut()
365            .ok_or_else(|| CLIConnectionError::new("transport is not connected"))?;
366
367        stdin
368            .write_all(data)
369            .await
370            .map_err(|e| CLIConnectionError::new(format!("failed to write to stdin: {}", e)))?;
371        stdin
372            .flush()
373            .await
374            .map_err(|e| CLIConnectionError::new(format!("failed to flush stdin: {}", e)))?;
375
376        Ok(())
377    }
378
379    async fn close_input(&mut self) -> Result<()> {
380        if let Some(mut stdin) = self.stdin.take() {
381            stdin
382                .shutdown()
383                .await
384                .map_err(|e| CLIConnectionError::new(format!("failed to close stdin: {}", e)))?;
385        }
386        Ok(())
387    }
388
389    async fn read(&mut self) -> Result<Option<Vec<u8>>> {
390        loop {
391            if let Some(data) = self.stdout_decoder.next() {
392                return Ok(Some(data));
393            }
394
395            let mut chunk = [0u8; 8192];
396            let read_result = {
397                let reader = self
398                    .stdout_reader
399                    .as_mut()
400                    .ok_or_else(|| CLIConnectionError::new("transport is not connected"))?;
401                reader.read(&mut chunk).await
402            };
403
404            match read_result {
405                Ok(0) => {
406                    self.stdout_decoder.finish()?;
407                    if let Some(data) = self.stdout_decoder.next() {
408                        return Ok(Some(data));
409                    }
410                    return self.finish_read().await;
411                }
412                Ok(n) => self
413                    .stdout_decoder
414                    .push(std::str::from_utf8(&chunk[..n]).map_err(|e| {
415                        CLIConnectionError::new(format!("stdout was not valid UTF-8: {}", e))
416                    })?)?,
417                Err(e) => {
418                    return Err(
419                        CLIConnectionError::new(format!("failed reading stdout: {}", e)).into(),
420                    )
421                }
422            }
423        }
424    }
425
426    async fn close(&mut self) -> Result<()> {
427        let _ = self.close_input().await;
428
429        if let Some(mut child) = self.child.take() {
430            let _ = child.kill().await;
431            let _ = child.wait().await;
432        }
433
434        Ok(())
435    }
436}
437
438#[cfg(test)]
439mod tests {
440    use super::*;
441
442    #[test]
443    fn process_env_matches_python_sdk_subprocess_defaults() {
444        let options = crate::types::ClaudeAgentOptions::builder()
445            .env_var("CLAUDE_CODE_ENTRYPOINT", "custom-entrypoint")
446            .env_var("TRACEPARENT", "explicit-trace")
447            .cwd("/tmp/project")
448            .enable_file_checkpointing(true)
449            .build();
450        let transport_options = TransportOptions::from(&options);
451
452        let env = build_process_env(
453            [
454                ("CLAUDECODE".to_string(), "1".to_string()),
455                ("PATH".to_string(), "/bin".to_string()),
456                ("TRACEPARENT".to_string(), "ambient-trace".to_string()),
457            ],
458            &transport_options,
459        );
460
461        assert_eq!(env.get("CLAUDECODE"), None);
462        assert_eq!(
463            env.get("CLAUDE_CODE_ENTRYPOINT").map(String::as_str),
464            Some("custom-entrypoint")
465        );
466        assert_eq!(
467            env.get("CLAUDE_AGENT_SDK_VERSION").map(String::as_str),
468            Some(env!("CARGO_PKG_VERSION"))
469        );
470        assert_eq!(
471            env.get("CLAUDE_CODE_ENABLE_SDK_FILE_CHECKPOINTING")
472                .map(String::as_str),
473            Some("true")
474        );
475        assert_eq!(env.get("PWD").map(String::as_str), Some("/tmp/project"));
476        assert_eq!(
477            env.get("TRACEPARENT").map(String::as_str),
478            Some("explicit-trace")
479        );
480    }
481
482    #[test]
483    fn process_env_injects_active_otel_context_like_python_sdk() {
484        let options = crate::types::ClaudeAgentOptions::builder().build();
485        let transport_options = TransportOptions::from(&options);
486        let mut env = build_process_env(
487            [
488                (
489                    "TRACEPARENT".to_string(),
490                    "00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01".to_string(),
491                ),
492                ("TRACESTATE".to_string(), "vendor=stale".to_string()),
493            ],
494            &transport_options,
495        );
496
497        apply_otel_trace_context(
498            &mut env,
499            &transport_options.env,
500            HashMap::from([
501                (
502                    "traceparent".to_string(),
503                    "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01".to_string(),
504                ),
505                ("tracestate".to_string(), "vendor=value".to_string()),
506            ]),
507        );
508
509        assert_eq!(
510            env.get("TRACEPARENT").map(String::as_str),
511            Some("00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01")
512        );
513        assert_eq!(
514            env.get("TRACESTATE").map(String::as_str),
515            Some("vendor=value")
516        );
517    }
518
519    #[test]
520    fn process_env_preserves_explicit_traceparent_over_otel_context() {
521        let options = crate::types::ClaudeAgentOptions::builder()
522            .env_var("TRACEPARENT", "custom")
523            .build();
524        let transport_options = TransportOptions::from(&options);
525        let mut env = build_process_env(
526            [("TRACEPARENT".to_string(), "ambient".to_string())],
527            &transport_options,
528        );
529
530        apply_otel_trace_context(
531            &mut env,
532            &transport_options.env,
533            HashMap::from([("traceparent".to_string(), "active".to_string())]),
534        );
535
536        assert_eq!(env.get("TRACEPARENT").map(String::as_str), Some("custom"));
537    }
538
539    #[test]
540    fn process_env_preserves_inherited_w3c_env_without_active_otel_span() {
541        let options = crate::types::ClaudeAgentOptions::builder().build();
542        let transport_options = TransportOptions::from(&options);
543        let mut env = build_process_env(
544            [
545                ("TRACEPARENT".to_string(), "ambient".to_string()),
546                ("TRACESTATE".to_string(), "vendor=abc".to_string()),
547            ],
548            &transport_options,
549        );
550
551        apply_otel_trace_context(
552            &mut env,
553            &transport_options.env,
554            HashMap::from([("baggage".to_string(), "user.id=123".to_string())]),
555        );
556
557        assert_eq!(env.get("TRACEPARENT").map(String::as_str), Some("ambient"));
558        assert_eq!(
559            env.get("TRACESTATE").map(String::as_str),
560            Some("vendor=abc")
561        );
562    }
563
564    #[tokio::test]
565    async fn subprocess_stderr_callback_receives_lines() {
566        use std::io::Write;
567        use std::sync::{Arc, Mutex};
568
569        let dir =
570            std::env::temp_dir().join(format!("claude-rust-stderr-test-{}", uuid::Uuid::new_v4()));
571        std::fs::create_dir_all(&dir).unwrap();
572        let script = dir.join("claude");
573        let mut file = std::fs::File::create(&script).unwrap();
574        writeln!(
575            file,
576            r#"#!/bin/sh
577if [ "$1" = "-v" ]; then
578  printf '2.0.0 (Claude Code)\n'
579  exit 0
580fi
581printf 'diagnostic line\n' >&2
582printf '{{"type":"result","subtype":"success","duration_ms":1,"duration_api_ms":1,"is_error":false,"num_turns":1,"session_id":"s"}}\n'
583"#
584        )
585        .unwrap();
586        #[cfg(unix)]
587        {
588            use std::os::unix::fs::PermissionsExt;
589            let mut permissions = std::fs::metadata(&script).unwrap().permissions();
590            permissions.set_mode(0o755);
591            std::fs::set_permissions(&script, permissions).unwrap();
592        }
593
594        let lines = Arc::new(Mutex::new(Vec::<String>::new()));
595        let captured = lines.clone();
596        let options = crate::types::ClaudeAgentOptions::builder()
597            .cli_path(script.to_string_lossy().to_string())
598            .stderr(move |line| captured.lock().unwrap().push(line))
599            .build();
600        let mut transport = SubprocessCLITransport::new(TransportOptions::from(&options));
601
602        transport.connect().await.unwrap();
603        let message = transport.read().await.unwrap().expect("result");
604        let value: serde_json::Value = serde_json::from_slice(&message).unwrap();
605        assert_eq!(value["type"], "result");
606
607        for _ in 0..20 {
608            if lines
609                .lock()
610                .unwrap()
611                .iter()
612                .any(|line| line == "diagnostic line\n")
613            {
614                let _ = transport.close().await;
615                let _ = std::fs::remove_dir_all(&dir);
616                return;
617            }
618            tokio::time::sleep(std::time::Duration::from_millis(10)).await;
619        }
620        panic!("stderr callback did not receive diagnostic line");
621    }
622}