Skip to main content

anthropic_sdk/streaming/
mod.rs

1//! Streaming support for real-time message generation.
2//!
3//! This module provides the `MessageStream` struct which handles Server-Sent Events (SSE)
4//! from the Anthropic API, accumulates messages from incremental updates, and provides
5//! an event-driven API for processing streaming responses.
6
7pub mod events;
8
9use std::collections::HashMap;
10use std::sync::{Arc, Mutex};
11use futures::Stream;
12use pin_project::pin_project;
13use tokio::sync::{broadcast, oneshot};
14use tokio_stream::wrappers::BroadcastStream;
15
16use crate::types::{
17    Message, MessageStreamEvent, ContentBlock, ContentBlockDelta, 
18    AnthropicError, Result
19};
20
21use self::events::{EventHandler, EventType};
22
23/// A streaming response from the Anthropic API.
24///
25/// `MessageStream` provides an event-driven interface for processing streaming responses
26/// from Claude. It accumulates message content from incremental updates and provides
27/// both callback-based and async iteration APIs.
28///
29/// # Examples
30///
31/// ## Callback-based processing:
32/// ```ignore
33/// # use anthropic_sdk::{Anthropic, MessageCreateBuilder};
34/// # async fn example() -> anthropic_sdk::Result<()> {
35/// let client = Anthropic::new("your-api-key")?;
36/// let stream = client.messages().create_stream(
37///     MessageCreateBuilder::new("claude-3-5-sonnet-latest", 1024)
38///         .user("Write a story about AI")
39///         .stream(true)
40///         .build()
41/// ).await?;
42///
43/// let final_message = stream
44///     .on_text(|delta, _snapshot| {
45///         print!("{}", delta);
46///     })
47///     .on_error(|error| {
48///         eprintln!("Stream error: {}", error);
49///     })
50///     .final_message().await?;
51/// # Ok(())
52/// # }
53/// ```
54///
55/// ## Async iteration:
56/// ```ignore
57/// # use anthropic_sdk::{Anthropic, MessageCreateBuilder, MessageStreamEvent};
58/// # use futures::StreamExt;
59/// # async fn example() -> anthropic_sdk::Result<()> {
60/// let client = Anthropic::new("your-api-key")?;
61/// let mut stream = client.messages().create_stream(
62///     MessageCreateBuilder::new("claude-3-5-sonnet-latest", 1024)
63///         .user("Tell me a joke")
64///         .stream(true)
65///         .build()
66/// ).await?;
67///
68/// while let Some(event) = stream.next().await {
69///     match event? {
70///         MessageStreamEvent::ContentBlockDelta { delta, .. } => {
71///             // Process incremental content
72///         }
73///         MessageStreamEvent::MessageStop => break,
74///         _ => {}
75///     }
76/// }
77/// # Ok(())
78/// # }
79/// ```
80#[pin_project]
81pub struct MessageStream {
82    /// Current accumulated message snapshot
83    current_message: Arc<Mutex<Option<Message>>>,
84    
85    /// Event handlers for different event types
86    event_handlers: Arc<Mutex<HashMap<EventType, Vec<EventHandler>>>>,
87    
88    /// Broadcast channel for distributing events to handlers
89    event_sender: broadcast::Sender<MessageStreamEvent>,
90    
91    /// Stream for events from the underlying HTTP stream
92    #[pin]
93    event_stream: BroadcastStream<MessageStreamEvent>,
94    
95    /// Channel for signaling when the stream ends
96    completion_sender: Option<oneshot::Sender<Result<Message>>>,
97    completion_receiver: oneshot::Receiver<Result<Message>>,
98    
99    /// Whether the stream has ended
100    ended: Arc<Mutex<bool>>,
101    
102    /// Whether an error occurred
103    errored: Arc<Mutex<bool>>,
104    
105    /// Whether the stream was aborted by the user
106    aborted: Arc<Mutex<bool>>,
107    
108    /// Response metadata
109    response: Option<reqwest::Response>,
110    request_id: Option<String>,
111}
112
113impl MessageStream {
114    /// Create a new MessageStream from an HTTP response.
115    ///
116    /// This is typically called internally by the SDK when creating streaming requests.
117    pub fn new(response: reqwest::Response, request_id: Option<String>) -> Self {
118        let (event_sender, event_receiver) = broadcast::channel(1000);
119        let (completion_sender, completion_receiver) = oneshot::channel();
120        
121        Self {
122            current_message: Arc::new(Mutex::new(None)),
123            event_handlers: Arc::new(Mutex::new(HashMap::new())),
124            event_sender,
125            event_stream: BroadcastStream::new(event_receiver),
126            completion_sender: Some(completion_sender),
127            completion_receiver,
128            ended: Arc::new(Mutex::new(false)),
129            errored: Arc::new(Mutex::new(false)),
130            aborted: Arc::new(Mutex::new(false)),
131            response: Some(response),
132            request_id,
133        }
134    }
135    
136    /// Create a new MessageStream from an HttpStreamClient.
137    ///
138    /// This connects a real HTTP stream to the MessageStream, providing
139    /// proper streaming functionality for real-time response processing.
140    pub fn from_http_stream(mut http_stream: crate::http::streaming::HttpStreamClient) -> Result<Self> {
141        let (event_sender, event_receiver) = broadcast::channel(1000);
142        let (completion_sender, completion_receiver) = oneshot::channel();
143        
144        let current_message = Arc::new(Mutex::new(None));
145        let ended = Arc::new(Mutex::new(false));
146        let errored = Arc::new(Mutex::new(false));
147        let request_id = http_stream.request_id().map(|s| s.to_string());
148        
149        // Clone references for the background task
150        let current_message_clone = current_message.clone();
151        let ended_clone = ended.clone();
152        let errored_clone = errored.clone();
153        let event_sender_clone = event_sender.clone();
154        
155        // Spawn task to process HTTP stream events
156        tokio::spawn(async move {
157            use futures::StreamExt;
158            let mut final_message: Option<crate::types::Message> = None;
159            
160            while let Some(event_result) = http_stream.next().await {
161                match event_result {
162                    Ok(event) => {
163                        // Update current message state
164                        match &event {
165                            crate::types::MessageStreamEvent::MessageStart { message } => {
166                                *current_message_clone.lock().unwrap() = Some(message.clone());
167                                final_message = Some(message.clone());
168                            }
169                            crate::types::MessageStreamEvent::ContentBlockStart { content_block, index } => {
170                                if let Some(ref mut msg) = *current_message_clone.lock().unwrap() {
171                                    while msg.content.len() <= *index {
172                                        msg.content.push(crate::types::ContentBlock::Text { text: String::new() });
173                                    }
174                                    msg.content[*index] = content_block.clone();
175                                }
176                                if let Some(ref mut msg) = final_message.as_mut() {
177                                    while msg.content.len() <= *index {
178                                        msg.content.push(crate::types::ContentBlock::Text { text: String::new() });
179                                    }
180                                    msg.content[*index] = content_block.clone();
181                                }
182                            }
183                            crate::types::MessageStreamEvent::ContentBlockDelta { delta, index } => {
184                                if let Some(ref mut msg) = *current_message_clone.lock().unwrap() {
185                                    if let Some(content_block) = msg.content.get_mut(*index) {
186                                        if let (crate::types::ContentBlock::Text { text }, 
187                                               crate::types::ContentBlockDelta::TextDelta { text: delta_text }) = 
188                                            (content_block, delta) {
189                                            text.push_str(delta_text);
190                                        }
191                                    }
192                                }
193                                if let Some(ref mut msg) = final_message.as_mut() {
194                                    if let Some(content_block) = msg.content.get_mut(*index) {
195                                        if let (crate::types::ContentBlock::Text { text }, 
196                                               crate::types::ContentBlockDelta::TextDelta { text: delta_text }) = 
197                                            (content_block, delta) {
198                                            text.push_str(delta_text);
199                                        }
200                                    }
201                                }
202                            }
203                            crate::types::MessageStreamEvent::MessageDelta { delta, usage } => {
204                                if let Some(ref mut msg) = *current_message_clone.lock().unwrap() {
205                                    if let Some(stop_reason) = &delta.stop_reason {
206                                        msg.stop_reason = Some(stop_reason.clone());
207                                    }
208                                    if let Some(stop_sequence) = &delta.stop_sequence {
209                                        msg.stop_sequence = Some(stop_sequence.clone());
210                                    }
211                                    msg.usage.output_tokens = usage.output_tokens;
212                                    if let Some(input_tokens) = usage.input_tokens {
213                                        msg.usage.input_tokens = input_tokens;
214                                    }
215                                }
216                                if let Some(ref mut msg) = final_message.as_mut() {
217                                    if let Some(stop_reason) = &delta.stop_reason {
218                                        msg.stop_reason = Some(stop_reason.clone());
219                                    }
220                                    if let Some(stop_sequence) = &delta.stop_sequence {
221                                        msg.stop_sequence = Some(stop_sequence.clone());
222                                    }
223                                    msg.usage.output_tokens = usage.output_tokens;
224                                    if let Some(input_tokens) = usage.input_tokens {
225                                        msg.usage.input_tokens = input_tokens;
226                                    }
227                                }
228                            }
229                            crate::types::MessageStreamEvent::MessageStop => {
230                                *ended_clone.lock().unwrap() = true;
231                                // Send the final message
232                                if let Some(message) = final_message.clone() {
233                                    let _ = completion_sender.send(Ok(message));
234                                } else {
235                                    let _ = completion_sender.send(Err(crate::types::AnthropicError::StreamError(
236                                        "Stream ended without message".to_string()
237                                    )));
238                                }
239                                // Send final event and break
240                                let _ = event_sender_clone.send(event);
241                                break;
242                            }
243                            _ => {}
244                        }
245                        
246                        // Send event to broadcast channel for callbacks
247                        let _ = event_sender_clone.send(event);
248                    }
249                    Err(e) => {
250                        *errored_clone.lock().unwrap() = true;
251                        let _ = completion_sender.send(Err(e));
252                        break;
253                    }
254                }
255            }
256        });
257        
258        Ok(Self {
259            current_message,
260            event_handlers: Arc::new(Mutex::new(HashMap::new())),
261            event_sender,
262            event_stream: BroadcastStream::new(event_receiver),
263            completion_sender: None, // Already consumed by the task
264            completion_receiver,
265            ended,
266            errored,
267            aborted: Arc::new(Mutex::new(false)),
268            response: None, // No response needed for HTTP stream
269            request_id,
270        })
271    }
272    
273    /// Register a callback for text delta events.
274    ///
275    /// The callback receives two parameters:
276    /// - `delta`: The new text being appended
277    /// - `snapshot`: The current accumulated text
278    ///
279    /// # Examples
280    /// ```rust,no_run
281    /// # use anthropic_sdk::MessageStream;
282    /// # async fn example(stream: MessageStream) {
283    /// stream.on_text(|delta, snapshot| {
284    ///     print!("{}", delta);
285    ///     println!("Total so far: {}", snapshot);
286    /// });
287    /// # }
288    /// ```
289    pub fn on_text<F>(self, callback: F) -> Self
290    where
291        F: Fn(&str, &str) + Send + Sync + 'static,
292    {
293        self.on(EventType::Text, EventHandler::Text(Box::new(callback)))
294    }
295    
296    /// Register a callback for stream events.
297    ///
298    /// This provides access to all raw stream events and the current message snapshot.
299    ///
300    /// # Examples
301    /// ```rust,no_run
302    /// # use anthropic_sdk::{MessageStream, MessageStreamEvent, Message};
303    /// # async fn example(stream: MessageStream) {
304    /// stream.on_stream_event(|event, snapshot| {
305    ///     match event {
306    ///         MessageStreamEvent::ContentBlockStart { .. } => {
307    ///             println!("New content block started");
308    ///         }
309    ///         _ => {}
310    ///     }
311    /// });
312    /// # }
313    /// ```
314    pub fn on_stream_event<F>(self, callback: F) -> Self
315    where
316        F: Fn(&MessageStreamEvent, &Message) + Send + Sync + 'static,
317    {
318        self.on(EventType::StreamEvent, EventHandler::StreamEvent(Box::new(callback)))
319    }
320    
321    /// Register a callback for when a complete message is received.
322    ///
323    /// # Examples
324    /// ```rust,no_run
325    /// # use anthropic_sdk::{MessageStream, Message};
326    /// # async fn example(stream: MessageStream) {
327    /// stream.on_message(|message| {
328    ///     println!("Received message: {:?}", message);
329    /// });
330    /// # }
331    /// ```
332    pub fn on_message<F>(self, callback: F) -> Self
333    where
334        F: Fn(&Message) + Send + Sync + 'static,
335    {
336        self.on(EventType::Message, EventHandler::Message(Box::new(callback)))
337    }
338    
339    /// Register a callback for when the final message is complete.
340    ///
341    /// # Examples
342    /// ```rust,no_run
343    /// # use anthropic_sdk::{MessageStream, Message};
344    /// # async fn example(stream: MessageStream) {
345    /// stream.on_final_message(|message| {
346    ///     println!("Final message: {:?}", message);
347    /// });
348    /// # }
349    /// ```
350    pub fn on_final_message<F>(self, callback: F) -> Self
351    where
352        F: Fn(&Message) + Send + Sync + 'static,
353    {
354        self.on(EventType::FinalMessage, EventHandler::FinalMessage(Box::new(callback)))
355    }
356    
357    /// Register a callback for errors.
358    ///
359    /// # Examples
360    /// ```rust,no_run
361    /// # use anthropic_sdk::{MessageStream, AnthropicError};
362    /// # async fn example(stream: MessageStream) {
363    /// stream.on_error(|error| {
364    ///     eprintln!("Stream error: {}", error);
365    /// });
366    /// # }
367    /// ```
368    pub fn on_error<F>(self, callback: F) -> Self
369    where
370        F: Fn(&AnthropicError) + Send + Sync + 'static,
371    {
372        self.on(EventType::Error, EventHandler::Error(Box::new(callback)))
373    }
374    
375    /// Register a callback for when the stream ends.
376    ///
377    /// # Examples
378    /// ```rust,no_run
379    /// # use anthropic_sdk::MessageStream;
380    /// # async fn example(stream: MessageStream) {
381    /// stream.on_end(|| {
382    ///     println!("Stream ended");
383    /// });
384    /// # }
385    /// ```
386    pub fn on_end<F>(self, callback: F) -> Self
387    where
388        F: Fn() + Send + Sync + 'static,
389    {
390        self.on(EventType::End, EventHandler::End(Box::new(callback)))
391    }
392    
393    /// Generic method to register event handlers.
394    fn on(self, event_type: EventType, handler: EventHandler) -> Self {
395        {
396            let mut handlers = self.event_handlers.lock().unwrap();
397            handlers.entry(event_type).or_insert_with(Vec::new).push(handler);
398        }
399        self
400    }
401    
402    /// Wait for the stream to complete and return the final message.
403    ///
404    /// This method will block until the stream ends and return the accumulated message.
405    ///
406    /// # Examples
407    /// ```rust,no_run
408    /// # use anthropic_sdk::MessageStream;
409    /// # async fn example(stream: MessageStream) -> anthropic_sdk::Result<()> {
410    /// let final_message = stream.final_message().await?;
411    /// println!("Claude said: {:?}", final_message.content);
412    /// # Ok(())
413    /// # }
414    /// ```
415    pub async fn final_message(self) -> Result<Message> {
416        self.completion_receiver.await
417            .map_err(|_| AnthropicError::StreamError("Stream ended unexpectedly".to_string()))?
418    }
419    
420    /// Wait for the stream to complete without returning the message.
421    ///
422    /// This is useful when you're processing events with callbacks and just need
423    /// to wait for completion.
424    ///
425    /// # Examples
426    /// ```rust,no_run
427    /// # use anthropic_sdk::MessageStream;
428    /// # async fn example(stream: MessageStream) -> anthropic_sdk::Result<()> {
429    /// stream.on_text(|delta, _| print!("{}", delta))
430    ///     .done().await?;
431    /// println!("\nStream completed!");
432    /// # Ok(())
433    /// # }
434    /// ```
435    pub async fn done(self) -> Result<()> {
436        self.completion_receiver.await
437            .map_err(|_| AnthropicError::StreamError("Stream ended unexpectedly".to_string()))?
438            .map(|_| ())
439    }
440    
441    /// Get the current accumulated message snapshot.
442    ///
443    /// Returns `None` if the stream hasn't started or no message has been received yet.
444    pub fn current_message(&self) -> Option<Message> {
445        self.current_message.lock().unwrap().clone()
446    }
447    
448    /// Check if the stream has ended.
449    pub fn ended(&self) -> bool {
450        *self.ended.lock().unwrap()
451    }
452    
453    /// Check if an error occurred.
454    pub fn errored(&self) -> bool {
455        *self.errored.lock().unwrap()
456    }
457    
458    /// Check if the stream was aborted.
459    pub fn aborted(&self) -> bool {
460        *self.aborted.lock().unwrap()
461    }
462    
463    /// Get the response metadata.
464    pub fn response(&self) -> Option<&reqwest::Response> {
465        self.response.as_ref()
466    }
467    
468    /// Get the request ID.
469    pub fn request_id(&self) -> Option<&str> {
470        self.request_id.as_deref()
471    }
472    
473    /// Abort the stream.
474    ///
475    /// This will cancel the underlying HTTP request and mark the stream as aborted.
476    pub fn abort(&self) {
477        *self.aborted.lock().unwrap() = true;
478        // In a real implementation, this would cancel the HTTP request
479    }
480    
481    /// Process a stream event and update the internal state.
482    ///
483    /// This method accumulates message content from incremental updates and
484    /// dispatches events to registered handlers.
485    #[allow(dead_code)]
486    fn process_event(&self, event: MessageStreamEvent) -> Result<()> {
487        // Update current message state based on the event
488        match &event {
489            MessageStreamEvent::MessageStart { message } => {
490                *self.current_message.lock().unwrap() = Some(message.clone());
491            }
492            MessageStreamEvent::ContentBlockStart { content_block, index } => {
493                if let Some(ref mut msg) = *self.current_message.lock().unwrap() {
494                    // Ensure the content array is large enough
495                    while msg.content.len() <= *index {
496                        msg.content.push(ContentBlock::Text { text: String::new() });
497                    }
498                    msg.content[*index] = content_block.clone();
499                }
500            }
501            MessageStreamEvent::ContentBlockDelta { delta, index } => {
502                if let Some(ref mut msg) = *self.current_message.lock().unwrap() {
503                    if let Some(content_block) = msg.content.get_mut(*index) {
504                        self.apply_delta(content_block, delta)?;
505                    }
506                }
507            }
508            MessageStreamEvent::MessageDelta { delta, usage } => {
509                if let Some(ref mut msg) = *self.current_message.lock().unwrap() {
510                    if let Some(stop_reason) = &delta.stop_reason {
511                        msg.stop_reason = Some(stop_reason.clone());
512                    }
513                    if let Some(stop_sequence) = &delta.stop_sequence {
514                        msg.stop_sequence = Some(stop_sequence.clone());
515                    }
516                    // Update usage information
517                    msg.usage.output_tokens = usage.output_tokens;
518                    if let Some(input_tokens) = usage.input_tokens {
519                        msg.usage.input_tokens = input_tokens;
520                    }
521                }
522            }
523            MessageStreamEvent::MessageStop => {
524                *self.ended.lock().unwrap() = true;
525            }
526            _ => {}
527        }
528        
529        // Dispatch event to handlers
530        self.dispatch_event(&event)?;
531        
532        // Send event to broadcast channel for async iteration
533        let _ = self.event_sender.send(event);
534        
535        Ok(())
536    }
537    
538    /// Apply a content block delta to update the content.
539    #[allow(dead_code)]
540    fn apply_delta(&self, content_block: &mut ContentBlock, delta: &ContentBlockDelta) -> Result<()> {
541        match (content_block, delta) {
542            (ContentBlock::Text { text }, ContentBlockDelta::TextDelta { text: delta_text }) => {
543                text.push_str(delta_text);
544            }
545            (ContentBlock::ToolUse { input, .. }, ContentBlockDelta::InputJsonDelta { partial_json }) => {
546                // In a real implementation, we'd parse the partial JSON
547                // For now, we'll just store it as-is
548                *input = serde_json::from_str(partial_json)
549                    .unwrap_or_else(|_| serde_json::Value::String(partial_json.clone()));
550            }
551            _ => {
552                // Other delta types would be handled here
553            }
554        }
555        Ok(())
556    }
557    
558    /// Dispatch an event to all registered handlers.
559    fn dispatch_event(&self, event: &MessageStreamEvent) -> Result<()> {
560        let handlers = self.event_handlers.lock().unwrap();
561        let current_message = self.current_message.lock().unwrap();
562        
563        // Dispatch to stream event handlers
564        if let Some(stream_handlers) = handlers.get(&EventType::StreamEvent) {
565            for handler in stream_handlers {
566                if let EventHandler::StreamEvent(callback) = handler {
567                    if let Some(ref msg) = *current_message {
568                        callback(event, msg);
569                    }
570                }
571            }
572        }
573        
574        // Dispatch specific event types
575        match event {
576            MessageStreamEvent::ContentBlockDelta { delta, .. } => {
577                if let ContentBlockDelta::TextDelta { text } = delta {
578                    if let Some(text_handlers) = handlers.get(&EventType::Text) {
579                        for handler in text_handlers {
580                            if let EventHandler::Text(callback) = handler {
581                                // Get current accumulated text for snapshot
582                                let snapshot = if let Some(ref msg) = *current_message {
583                                    self.get_accumulated_text(msg)
584                                } else {
585                                    String::new()
586                                };
587                                callback(text, &snapshot);
588                            }
589                        }
590                    }
591                }
592            }
593            MessageStreamEvent::MessageStop => {
594                if let Some(end_handlers) = handlers.get(&EventType::End) {
595                    for handler in end_handlers {
596                        if let EventHandler::End(callback) = handler {
597                            callback();
598                        }
599                    }
600                }
601                
602                // Send final message
603                if let Some(ref msg) = *current_message {
604                    if let Some(final_handlers) = handlers.get(&EventType::FinalMessage) {
605                        for handler in final_handlers {
606                            if let EventHandler::FinalMessage(callback) = handler {
607                                callback(msg);
608                            }
609                        }
610                    }
611                }
612            }
613            _ => {}
614        }
615        
616        Ok(())
617    }
618    
619    /// Get the accumulated text from all text content blocks.
620    fn get_accumulated_text(&self, message: &Message) -> String {
621        message.content
622            .iter()
623            .filter_map(|block| match block {
624                ContentBlock::Text { text } => Some(text.as_str()),
625                _ => None,
626            })
627            .collect::<Vec<_>>()
628            .join("")
629    }
630}
631
632impl Stream for MessageStream {
633    type Item = Result<MessageStreamEvent>;
634    
635    fn poll_next(
636        self: std::pin::Pin<&mut Self>,
637        cx: &mut std::task::Context<'_>,
638    ) -> std::task::Poll<Option<Self::Item>> {
639        use futures::Stream as FuturesStream;
640        
641        let this = self.project();
642        
643        match FuturesStream::poll_next(this.event_stream, cx) {
644            std::task::Poll::Ready(Some(Ok(event))) => {
645                std::task::Poll::Ready(Some(Ok(event)))
646            }
647            std::task::Poll::Ready(Some(Err(err))) => {
648                // Handle any broadcast stream errors
649                std::task::Poll::Ready(Some(Err(AnthropicError::StreamError(
650                    format!("Stream error: {}", err)
651                ))))
652            }
653            std::task::Poll::Ready(None) => {
654                std::task::Poll::Ready(None)
655            }
656            std::task::Poll::Pending => std::task::Poll::Pending,
657        }
658    }
659}
660
661#[cfg(test)]
662mod tests {
663    use super::*;
664    use crate::types::{Role, Usage};
665    
666    // For testing, we'll use a simple helper to create a dummy response
667    async fn create_dummy_response() -> reqwest::Response {
668        // Create a simple HTTP client and make a basic request for testing
669        let client = reqwest::Client::new();
670        // Use httpbin.org which provides testing endpoints
671        client.get("https://httpbin.org/status/200")
672            .send()
673            .await
674            .expect("Failed to create test response")
675    }
676    
677    #[tokio::test]
678    async fn test_message_stream_creation() {
679        let response = create_dummy_response().await;
680        let stream = MessageStream::new(response, Some("test-request-id".to_string()));
681        
682        assert!(!stream.ended());
683        assert!(!stream.errored());
684        assert!(!stream.aborted());
685        assert_eq!(stream.request_id(), Some("test-request-id"));
686    }
687    
688    #[tokio::test]
689    async fn test_event_processing() {
690        let response = create_dummy_response().await;
691        let stream = MessageStream::new(response, None);
692        
693        // Test message start event
694        let start_event = MessageStreamEvent::MessageStart {
695            message: Message {
696                id: "msg_test".to_string(),
697                type_: "message".to_string(),
698                role: Role::Assistant,
699                content: vec![],
700                model: "claude-3-5-sonnet-latest".to_string(),
701                stop_reason: None,
702                stop_sequence: None,
703                usage: Usage {
704                    input_tokens: 10,
705                    output_tokens: 0,
706                    cache_creation_input_tokens: None,
707                    cache_read_input_tokens: None,
708                    server_tool_use: None,
709                    service_tier: None,
710                },
711                request_id: None,
712            },
713        };
714        
715        stream.process_event(start_event).unwrap();
716        
717        let current = stream.current_message().unwrap();
718        assert_eq!(current.id, "msg_test");
719        assert_eq!(current.role, Role::Assistant);
720    }
721    
722    #[test]
723    fn test_event_handlers() {
724        use std::sync::{Arc, Mutex};
725        use std::collections::HashMap;
726        
727        // Test creating event handlers directly
728        let text_called = Arc::new(Mutex::new(false));
729        let text_called_clone = text_called.clone();
730        
731        let _handler = EventHandler::Text(Box::new(move |_delta, _snapshot| {
732            *text_called_clone.lock().unwrap() = true;
733        }));
734        
735        // Test event type equality
736        assert_eq!(EventType::Text, EventType::Text);
737        assert_ne!(EventType::Text, EventType::Error);
738        
739        // Test using event types as hash keys
740        let mut map: HashMap<EventType, String> = HashMap::new();
741        map.insert(EventType::Text, "text_handler".to_string());
742        assert_eq!(map.get(&EventType::Text), Some(&"text_handler".to_string()));
743    }
744}