1use futures::{Stream, StreamExt};
4
5pub struct SpawnedStream {
12 pub events: mpsc::UnboundedReceiver<StreamEvent>,
13 pub cleanup: SpawnedStreamCleanup,
14}
15
16#[derive(Clone)]
17pub struct SpawnedStreamCleanup {
18 task: Arc<Mutex<Option<tokio::task::JoinHandle<Result<()>>>>>,
19}
20
21impl SpawnedStreamCleanup {
22 pub async fn wait_for_cleanup(&self) -> Result<()> {
26 let task = self.task.lock().await.take();
27 let Some(task) = task else {
28 return Ok(());
29 };
30 match task.await {
31 Ok(result) => result,
32 Err(error) => Err(crate::error::ClaudeSDKError::Other(format!(
33 "spawned Claude stream task did not finish cleanly: {error}"
34 ))),
35 }
36 }
37}
38use std::collections::HashMap;
39use std::sync::Arc;
40use tokio::sync::{mpsc, Mutex, RwLock};
41
42use crate::client_stream::stream_events_from_message;
43use crate::client_types::{MessageResponse, StreamEvent};
44use crate::error::{CLIConnectionError, Result};
45use crate::internal::control::{
46 initialize_request, initialize_timeout_duration, respond_to_control_request,
47 send_control_request_with_callbacks, send_control_request_with_callbacks_and_timeout,
48 ControlCallbacks,
49};
50use crate::internal::parser::parse_message_line;
51use crate::internal::session_resume::{
52 apply_materialized_options, materialize_resume_session, MaterializedResume,
53};
54use crate::internal::session_store_validation::validate_session_store_options;
55use crate::internal::transcript_mirror::TranscriptMirrorBatcher;
56use crate::internal::transport::{SubprocessCLITransport, Transport, TransportOptions};
57use crate::types::{
58 ClaudeAgentOptions, ContentBlock, ContextUsageResponse, MCPStatusResponse, Message,
59 PermissionMode, UserMessageInput,
60};
61
62#[derive(Debug)]
63#[allow(dead_code)]
64struct ClientState {
65 messages: Vec<Message>,
66 current_stream_buffer: String,
67 is_streaming: bool,
68 server_info: Option<HashMap<String, serde_json::Value>>,
69}
70
71pub struct ClaudeAgentClient {
72 transport: Box<dyn Transport>,
73 state: Arc<RwLock<ClientState>>,
74 session_id: String,
75 connected: bool,
76 initialized: bool,
77 initialization_result: Option<serde_json::Map<String, serde_json::Value>>,
78 control_callbacks: ControlCallbacks,
79 transcript_mirror: Option<TranscriptMirrorBatcher>,
80 source_options: Option<ClaudeAgentOptions>,
81 materialized_resume: Option<MaterializedResume>,
82}
83
84impl std::fmt::Debug for ClaudeAgentClient {
85 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
86 f.debug_struct("ClaudeAgentClient")
87 .field("session_id", &self.session_id)
88 .finish_non_exhaustive()
89 }
90}
91
92impl ClaudeAgentClient {
93 pub fn spawn_stream_message(
107 options: ClaudeAgentOptions,
108 content: impl Into<UserMessageInput>,
109 ) -> mpsc::UnboundedReceiver<StreamEvent> {
110 Self::spawn_stream_message_supervised(options, content).events
111 }
112
113 pub fn spawn_stream_message_supervised(
116 options: ClaudeAgentOptions,
117 content: impl Into<UserMessageInput>,
118 ) -> SpawnedStream {
119 let content = content.into();
120 let (tx, rx) = mpsc::unbounded_channel();
121 let task = tokio::spawn(async move {
122 let result = Self::run_spawned_stream(options, content, tx.clone()).await;
123 if let Err(ref err) = result {
124 if !tx.is_closed() {
125 let _ = tx.send(StreamEvent::Error(err.to_string()));
126 }
127 }
128 result
129 });
130 SpawnedStream {
131 events: rx,
132 cleanup: SpawnedStreamCleanup {
133 task: Arc::new(Mutex::new(Some(task))),
134 },
135 }
136 }
137
138 async fn run_spawned_stream(
139 options: ClaudeAgentOptions,
140 content: UserMessageInput,
141 tx: mpsc::UnboundedSender<StreamEvent>,
142 ) -> Result<()> {
143 let client = Self::new(options)?;
144 Self::run_client_stream(client, content, tx).await
145 }
146
147 pub async fn run_client_stream(
151 mut client: Self,
152 content: UserMessageInput,
153 tx: mpsc::UnboundedSender<StreamEvent>,
154 ) -> Result<()> {
155 let result = async {
156 client.connect().await?;
157 client.require_connected()?;
158 let payload = client.build_user_payload(&content, None)?;
159 let json_payload = serde_json::to_vec(&payload)?;
160 client.transport.write(&json_payload).await?;
161 client.transport.write(b"\n").await?;
162 {
163 let mut state = client.state.write().await;
164 state.is_streaming = true;
165 }
166 loop {
167 let data = tokio::select! {
168 _ = tx.closed() => break,
169 result = client.transport.read() => result?,
170 };
171 let Some(data) = data else {
172 break;
173 };
174 let line = String::from_utf8_lossy(&data);
175 let value = serde_json::from_slice::<serde_json::Value>(&data)?;
176 if value.get("type").and_then(|v| v.as_str()) == Some("control_request") {
177 respond_to_control_request(
178 client.transport.as_mut(),
179 &value,
180 &client.control_callbacks,
181 )
182 .await?;
183 continue;
184 }
185 if value.get("type").and_then(|v| v.as_str()) == Some("transcript_mirror") {
186 if let Some(batcher) = &mut client.transcript_mirror {
187 for message in batcher.enqueue_value(&value).await? {
188 let _ = tx.send(StreamEvent::Error(format!("{message:?}")));
189 }
190 }
191 continue;
192 }
193 let message = match parse_message_line(&line) {
194 Ok(Some(message)) => message,
195 Ok(None) => continue,
196 Err(err) => {
197 tracing::warn!("skipping unparseable CLI message: {err}");
200 continue;
201 }
202 };
203 for event in stream_events_from_message(&message, &client.session_id) {
204 let _ = tx.send(event);
205 }
206 let done = matches!(message, Message::ResultMsg { .. });
207 if done {
208 if let Some(batcher) = &mut client.transcript_mirror {
209 for message in batcher.flush().await? {
210 let _ = tx.send(StreamEvent::Error(format!("{message:?}")));
211 }
212 }
213 }
214 {
215 let mut state = client.state.write().await;
216 state.messages.push(message);
217 if done {
218 state.is_streaming = false;
219 }
220 }
221 if done {
222 break;
223 }
224 }
225 Ok(())
226 }
227 .await;
228 let close_result = client.disconnect().await;
232 result.and(close_result)
233 }
234
235 pub fn new(options: ClaudeAgentOptions) -> Result<Self> {
236 validate_session_store_options(&options)?;
237 let transport_options = TransportOptions::from(&options);
238 let transport = SubprocessCLITransport::new(transport_options);
239 let mut client = Self::with_transport(options.clone(), Box::new(transport))?;
240 client.source_options = Some(options);
241 Ok(client)
242 }
243
244 pub fn with_transport(
245 options: ClaudeAgentOptions,
246 transport: Box<dyn Transport>,
247 ) -> Result<Self> {
248 let session_id = options
249 .session_id
250 .clone()
251 .or_else(|| options.resume.clone())
252 .unwrap_or_else(|| "default".to_string());
253 let state = Arc::new(RwLock::new(ClientState {
254 messages: Vec::new(),
255 current_stream_buffer: String::new(),
256 is_streaming: false,
257 server_info: None,
258 }));
259 Ok(Self {
260 transport,
261 state,
262 session_id,
263 connected: false,
264 initialized: false,
265 initialization_result: None,
266 control_callbacks: ControlCallbacks::from_options(&options),
267 transcript_mirror: TranscriptMirrorBatcher::from_options(&options),
268 source_options: None,
269 materialized_resume: None,
270 })
271 }
272
273 pub async fn connect(&mut self) -> Result<()> {
274 if !self.connected {
275 self.materialize_resume_before_connect().await?;
276 self.transport.connect().await?;
277 self.connected = true;
278 }
279 self.ensure_initialized().await?;
280 Ok(())
281 }
282
283 pub async fn connect_with_prompt(
284 &mut self,
285 content: impl Into<UserMessageInput>,
286 ) -> Result<()> {
287 self.connect().await?;
288 let content = content.into();
289 let payload = self.build_user_payload(&content, None)?;
290 let mut json_payload = serde_json::to_vec(&payload)?;
291 json_payload.push(b'\n');
292 self.transport.write(&json_payload).await
293 }
294
295 pub async fn connect_with_stream<S>(&mut self, stream: S) -> Result<()>
296 where
297 S: Stream<Item = serde_json::Value> + Unpin,
298 {
299 self.connect().await?;
300 self.write_message_stream(stream, "default").await
301 }
302
303 async fn materialize_resume_before_connect(&mut self) -> Result<()> {
304 let Some(options) = self.source_options.clone() else {
305 return Ok(());
306 };
307 let Some(materialized) = materialize_resume_session(&options).await? else {
308 return Ok(());
309 };
310 let options = apply_materialized_options(&options, &materialized);
311 self.session_id = options
312 .session_id
313 .clone()
314 .or_else(|| options.resume.clone())
315 .unwrap_or_else(|| "default".to_string());
316 self.transport = Box::new(SubprocessCLITransport::new(TransportOptions::from(
317 &options,
318 )));
319 self.transcript_mirror = TranscriptMirrorBatcher::from_options(&options);
320 self.source_options = Some(options);
321 self.materialized_resume = Some(materialized);
322 Ok(())
323 }
324
325 fn require_connected(&self) -> Result<()> {
326 if self.connected && self.initialized {
327 Ok(())
328 } else {
329 Err(CLIConnectionError::new("Not connected. Call connect() first.").into())
330 }
331 }
332
333 async fn ensure_initialized(&mut self) -> Result<()> {
334 if self.initialized {
335 return Ok(());
336 }
337
338 let response = send_control_request_with_callbacks_and_timeout(
339 self.transport.as_mut(),
340 initialize_request(&self.control_callbacks),
341 &self.control_callbacks,
342 initialize_timeout_duration(),
343 )
344 .await?;
345 self.initialization_result = Some(response);
346 self.initialized = true;
347 Ok(())
348 }
349
350 pub async fn send_message(
351 &mut self,
352 content: impl Into<UserMessageInput>,
353 ) -> Result<MessageResponse> {
354 self.query(content).await?;
355 let messages = self.receive_response().await?;
356 let mut content_parts: Vec<String> = Vec::new();
357 let mut blocks: Vec<ContentBlock> = Vec::new();
358 let mut usage: Option<HashMap<String, serde_json::Value>> = None;
359 let mut stop_reason: Option<String> = None;
360 let mut model = String::new();
361
362 for message in messages {
363 match message {
364 Message::AssistantMsg {
365 content: assistant_content,
366 ..
367 } => {
368 if model.is_empty() {
370 model.clone_from(&assistant_content.model);
371 }
372 for block in &assistant_content.content {
373 match block {
374 ContentBlock::Text { text } => content_parts.push(text.clone()),
375 ContentBlock::Thinking { thinking, .. } => {
376 content_parts.push(thinking.clone())
377 }
378 _ => {}
379 }
380 blocks.push(block.clone());
381 }
382 }
383 Message::ResultMsg {
384 stop_reason: reason,
385 usage: u,
386 ..
387 } => {
388 stop_reason = reason;
389 if let Some(u) = u {
390 usage = Some(u.into_iter().collect());
391 }
392 }
393 _ => {}
394 }
395 }
396
397 Ok(MessageResponse {
398 content: content_parts.join(""),
399 blocks,
400 model,
401 stop_reason,
402 session_id: self.session_id.clone(),
403 usage,
404 })
405 }
406
407 pub async fn query(&mut self, content: impl Into<UserMessageInput>) -> Result<()> {
408 self.require_connected()?;
409 let content = content.into();
410 let payload = self.build_user_payload(&content, None)?;
411 let mut json_payload = serde_json::to_vec(&payload)?;
412 json_payload.push(b'\n');
413 self.transport.write(&json_payload).await
414 }
415
416 pub async fn query_with_session_id(
417 &mut self,
418 content: impl Into<UserMessageInput>,
419 session_id: impl Into<String>,
420 ) -> Result<()> {
421 self.require_connected()?;
422 let content = content.into();
423 let session_id = session_id.into();
424 let payload = self.build_user_payload(&content, Some(&session_id))?;
425 let mut json_payload = serde_json::to_vec(&payload)?;
426 json_payload.push(b'\n');
427 self.transport.write(&json_payload).await
428 }
429
430 pub async fn query_stream<S>(&mut self, stream: S) -> Result<()>
431 where
432 S: Stream<Item = serde_json::Value> + Unpin,
433 {
434 self.query_stream_with_session_id(stream, "default").await
435 }
436
437 pub async fn query_stream_with_session_id<S>(
438 &mut self,
439 stream: S,
440 session_id: impl Into<String>,
441 ) -> Result<()>
442 where
443 S: Stream<Item = serde_json::Value> + Unpin,
444 {
445 self.require_connected()?;
446 self.write_message_stream(stream, &session_id.into()).await
447 }
448
449 pub async fn receive_response(&mut self) -> Result<Vec<Message>> {
450 self.receive_messages_until(true).await
451 }
452
453 pub async fn receive_messages(&mut self) -> Result<Vec<Message>> {
454 self.receive_messages_until(false).await
455 }
456
457 async fn receive_messages_until(&mut self, stop_at_result: bool) -> Result<Vec<Message>> {
458 self.require_connected()?;
459 let mut messages = Vec::new();
460 while let Some(data) = self.transport.read().await? {
461 let line = String::from_utf8_lossy(&data);
462 let value = serde_json::from_slice::<serde_json::Value>(&data)?;
463 if value.get("type").and_then(|v| v.as_str()) == Some("control_request") {
464 respond_to_control_request(
465 self.transport.as_mut(),
466 &value,
467 &self.control_callbacks,
468 )
469 .await?;
470 continue;
471 }
472 if value.get("type").and_then(|v| v.as_str()) == Some("transcript_mirror") {
473 if let Some(batcher) = &mut self.transcript_mirror {
474 messages.extend(batcher.enqueue_value(&value).await?);
475 }
476 continue;
477 }
478 let message = match parse_message_line(&line) {
479 Ok(Some(message)) => message,
480 Ok(None) => continue,
481 Err(err) => {
482 tracing::warn!("skipping unparseable CLI message: {err}");
485 continue;
486 }
487 };
488 let done = matches!(message, Message::ResultMsg { .. });
489 if done {
490 if let Some(batcher) = &mut self.transcript_mirror {
491 messages.extend(batcher.flush().await?);
492 }
493 }
494 {
495 let mut state = self.state.write().await;
496 state.messages.push(message.clone());
497 }
498 messages.push(message);
499 if stop_at_result && done {
500 break;
501 }
502 }
503 Ok(messages)
504 }
505
506 pub async fn stream_message(
507 &mut self,
508 content: impl Into<UserMessageInput>,
509 ) -> Result<mpsc::UnboundedReceiver<StreamEvent>> {
510 self.require_connected()?;
511 let content = content.into();
512 let payload = self.build_user_payload(&content, None)?;
513 let json_payload = serde_json::to_vec(&payload)?;
514 self.transport.write(&json_payload).await?;
515 self.transport
516 .write(
517 b"
518",
519 )
520 .await?;
521 let (tx, rx) = mpsc::unbounded_channel();
522 {
523 let mut state = self.state.write().await;
524 state.is_streaming = true;
525 }
526 while let Some(data) = self.transport.read().await? {
527 let line = String::from_utf8_lossy(&data);
528 let value = serde_json::from_slice::<serde_json::Value>(&data)?;
529 if value.get("type").and_then(|v| v.as_str()) == Some("control_request") {
530 respond_to_control_request(
531 self.transport.as_mut(),
532 &value,
533 &self.control_callbacks,
534 )
535 .await?;
536 continue;
537 }
538 if value.get("type").and_then(|v| v.as_str()) == Some("transcript_mirror") {
539 if let Some(batcher) = &mut self.transcript_mirror {
540 for message in batcher.enqueue_value(&value).await? {
541 let _ = tx.send(StreamEvent::Error(format!("{message:?}")));
542 }
543 }
544 continue;
545 }
546 let message = match parse_message_line(&line) {
547 Ok(Some(message)) => message,
548 Ok(None) => continue,
549 Err(err) => {
550 tracing::warn!("skipping unparseable CLI message: {err}");
553 continue;
554 }
555 };
556 for event in stream_events_from_message(&message, &self.session_id) {
557 let _ = tx.send(event);
558 }
559 let done = matches!(message, Message::ResultMsg { .. });
560 if done {
561 if let Some(batcher) = &mut self.transcript_mirror {
562 for message in batcher.flush().await? {
563 let _ = tx.send(StreamEvent::Error(format!("{message:?}")));
564 }
565 }
566 }
567 {
568 let mut state = self.state.write().await;
569 state.messages.push(message);
570 if done {
571 state.is_streaming = false;
572 }
573 }
574 if done {
575 break;
576 }
577 }
578 Ok(rx)
579 }
580
581 async fn write_message_stream<S>(&mut self, mut stream: S, session_id: &str) -> Result<()>
582 where
583 S: Stream<Item = serde_json::Value> + Unpin,
584 {
585 while let Some(mut message) = stream.next().await {
586 if let Some(object) = message.as_object_mut() {
587 object
588 .entry("session_id")
589 .or_insert_with(|| serde_json::Value::String(session_id.to_string()));
590 }
591 let mut json_payload = serde_json::to_vec(&message)?;
592 json_payload.push(b'\n');
593 self.transport.write(&json_payload).await?;
594 }
595 Ok(())
596 }
597
598 pub async fn get_conversation_history(&self) -> Result<Vec<Message>> {
599 let state = self.state.read().await;
600 Ok(state.messages.clone())
601 }
602
603 pub async fn abort(&mut self) -> Result<()> {
604 if let Some(batcher) = &mut self.transcript_mirror {
605 let _ = batcher.flush().await?;
606 }
607 self.transport.close().await?;
608 if let Some(materialized) = &self.materialized_resume {
609 materialized.cleanup().await;
610 }
611 self.materialized_resume = None;
612 self.connected = false;
613 self.initialized = false;
614 Ok(())
615 }
616
617 pub async fn disconnect(&mut self) -> Result<()> {
618 self.abort().await
619 }
620
621 pub async fn close(mut self) -> Result<()> {
622 if let Some(batcher) = &mut self.transcript_mirror {
623 let _ = batcher.flush().await?;
624 }
625 self.transport.close().await?;
626 if let Some(materialized) = &self.materialized_resume {
627 materialized.cleanup().await;
628 }
629 Ok(())
630 }
631
632 pub async fn interrupt(&mut self) -> Result<()> {
633 self.require_connected()?;
634 send_control_request_with_callbacks(
635 self.transport.as_mut(),
636 serde_json::json!({"subtype": "interrupt"}),
637 &self.control_callbacks,
638 )
639 .await?;
640 Ok(())
641 }
642
643 pub async fn set_permission_mode(&mut self, mode: PermissionMode) -> Result<()> {
644 self.require_connected()?;
645 send_control_request_with_callbacks(
646 self.transport.as_mut(),
647 serde_json::json!({
648 "subtype": "set_permission_mode",
649 "mode": mode,
650 }),
651 &self.control_callbacks,
652 )
653 .await?;
654 Ok(())
655 }
656
657 pub async fn set_model(&mut self, model: Option<String>) -> Result<()> {
658 self.require_connected()?;
659 let model = model.map(serde_json::Value::String);
660 send_control_request_with_callbacks(
661 self.transport.as_mut(),
662 serde_json::json!({
663 "subtype": "set_model",
664 "model": model.unwrap_or(serde_json::Value::Null),
665 }),
666 &self.control_callbacks,
667 )
668 .await?;
669 Ok(())
670 }
671
672 pub async fn rewind_files(&mut self, user_message_id: impl Into<String>) -> Result<()> {
673 self.require_connected()?;
674 send_control_request_with_callbacks(
675 self.transport.as_mut(),
676 serde_json::json!({
677 "subtype": "rewind_files",
678 "user_message_id": user_message_id.into(),
679 }),
680 &self.control_callbacks,
681 )
682 .await?;
683 Ok(())
684 }
685
686 pub async fn reconnect_mcp_server(&mut self, server_name: impl Into<String>) -> Result<()> {
687 self.require_connected()?;
688 send_control_request_with_callbacks(
689 self.transport.as_mut(),
690 serde_json::json!({
691 "subtype": "mcp_reconnect",
692 "serverName": server_name.into(),
693 }),
694 &self.control_callbacks,
695 )
696 .await?;
697 Ok(())
698 }
699
700 pub async fn toggle_mcp_server(
701 &mut self,
702 server_name: impl Into<String>,
703 enabled: bool,
704 ) -> Result<()> {
705 self.require_connected()?;
706 send_control_request_with_callbacks(
707 self.transport.as_mut(),
708 serde_json::json!({
709 "subtype": "mcp_toggle",
710 "serverName": server_name.into(),
711 "enabled": enabled,
712 }),
713 &self.control_callbacks,
714 )
715 .await?;
716 Ok(())
717 }
718
719 pub async fn stop_task(&mut self, task_id: impl Into<String>) -> Result<()> {
720 self.require_connected()?;
721 send_control_request_with_callbacks(
722 self.transport.as_mut(),
723 serde_json::json!({
724 "subtype": "stop_task",
725 "task_id": task_id.into(),
726 }),
727 &self.control_callbacks,
728 )
729 .await?;
730 Ok(())
731 }
732
733 pub async fn get_mcp_status(&mut self) -> Result<MCPStatusResponse> {
734 self.require_connected()?;
735 let response = send_control_request_with_callbacks(
736 self.transport.as_mut(),
737 serde_json::json!({"subtype": "mcp_status"}),
738 &self.control_callbacks,
739 )
740 .await?;
741 let value = serde_json::Value::Object(response);
742 Ok(serde_json::from_value(value)?)
743 }
744
745 pub async fn get_context_usage(&mut self) -> Result<ContextUsageResponse> {
746 self.require_connected()?;
747 let response = send_control_request_with_callbacks(
748 self.transport.as_mut(),
749 serde_json::json!({"subtype": "get_context_usage"}),
750 &self.control_callbacks,
751 )
752 .await?;
753 Ok(serde_json::from_value(serde_json::Value::Object(response))?)
754 }
755
756 pub fn get_server_info(&self) -> Option<&serde_json::Map<String, serde_json::Value>> {
757 self.initialization_result.as_ref()
758 }
759
760 fn build_user_payload(
761 &self,
762 content: &UserMessageInput,
763 session_id: Option<&str>,
764 ) -> Result<serde_json::Map<String, serde_json::Value>> {
765 let mut payload = serde_json::Map::new();
766 payload.insert(
767 "type".to_string(),
768 serde_json::Value::String("user".to_string()),
769 );
770 payload.insert(
771 "session_id".to_string(),
772 serde_json::Value::String(
773 session_id
774 .map(String::from)
775 .unwrap_or_else(|| self.session_id.clone()),
776 ),
777 );
778 let message = serde_json::json!({"role": "user", "content": content.to_content_value()});
781 payload.insert("message".to_string(), message);
782 Ok(payload)
783 }
784}