Skip to main content

llm_optimizer_integrations/anthropic/
streaming.rs

1//! Anthropic Claude API streaming support
2//!
3//! Handles Server-Sent Events (SSE) streaming from Claude API.
4
5use super::types::*;
6use anyhow::{anyhow, Context, Result};
7use futures::stream::{Stream, StreamExt};
8use reqwest::header::{HeaderMap, HeaderValue, CONTENT_TYPE};
9use std::pin::Pin;
10use std::sync::Arc;
11use tokio::sync::RwLock;
12use tracing::{debug, error, info, warn};
13
14/// Stream handler for Claude API streaming responses
15pub struct StreamHandler {
16    /// Configuration
17    config: Arc<RwLock<AnthropicConfig>>,
18    /// HTTP client
19    client: reqwest::Client,
20    /// Cost tracker
21    cost_tracker: Arc<RwLock<CostTracker>>,
22}
23
24impl StreamHandler {
25    /// Create a new stream handler
26    ///
27    /// # Arguments
28    ///
29    /// * `config` - API configuration
30    /// * `client` - HTTP client to use
31    /// * `cost_tracker` - Cost tracker to update
32    pub fn new(
33        config: Arc<RwLock<AnthropicConfig>>,
34        client: reqwest::Client,
35        cost_tracker: Arc<RwLock<CostTracker>>,
36    ) -> Self {
37        Self {
38            config,
39            client,
40            cost_tracker,
41        }
42    }
43
44    /// Send a streaming message request
45    ///
46    /// # Arguments
47    ///
48    /// * `request` - Message request with stream=true
49    ///
50    /// # Returns
51    ///
52    /// Returns a stream of events
53    pub async fn stream_message(
54        &self,
55        mut request: MessageRequest,
56    ) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent>> + Send>>> {
57        // Ensure streaming is enabled
58        request.stream = true;
59
60        let config = self.config.read().await;
61        let url = format!("{}/v1/messages", config.base_url);
62
63        let headers = self.build_headers(&config)?;
64
65        debug!(
66            "Starting streaming request to model: {}",
67            request.model
68        );
69
70        let response = self
71            .client
72            .post(&url)
73            .headers(headers)
74            .json(&request)
75            .send()
76            .await
77            .context("Failed to send streaming request")?;
78
79        if !response.status().is_success() {
80            let status = response.status();
81            let error_text = response.text().await.unwrap_or_default();
82            error!("Streaming request failed ({}): {}", status, error_text);
83            return Err(anyhow!("Streaming request failed: {}", error_text));
84        }
85
86        info!("Streaming response started");
87
88        // Parse SSE stream - create stream from response chunks
89        let byte_stream = futures::stream::unfold(response, |mut resp| async {
90            match resp.chunk().await {
91                Ok(Some(chunk)) => Some((Ok(chunk), resp)),
92                Ok(None) => None,
93                Err(e) => Some((Err(e), resp)),
94            }
95        });
96
97        let stream = self.parse_sse_stream(byte_stream);
98
99        Ok(Box::pin(stream))
100    }
101
102    /// Complete a streaming request and collect all text
103    ///
104    /// # Arguments
105    ///
106    /// * `request` - Message request
107    ///
108    /// # Returns
109    ///
110    /// Returns the complete text response and usage stats
111    pub async fn stream_complete(
112        &self,
113        request: MessageRequest,
114    ) -> Result<(String, Usage)> {
115        let mut stream = self.stream_message(request).await?;
116
117        let mut text = String::new();
118        let mut usage = Usage {
119            input_tokens: 0,
120            output_tokens: 0,
121        };
122
123        while let Some(event_result) = stream.next().await {
124            match event_result? {
125                StreamEvent::ContentBlockDelta { delta, .. } => {
126                    if let Delta::TextDelta { text: delta_text } = delta {
127                        text.push_str(&delta_text);
128                    }
129                }
130                StreamEvent::MessageStart { message } => {
131                    usage = message.usage;
132                }
133                StreamEvent::MessageDelta {
134                    usage: final_usage,
135                    ..
136                } => {
137                    usage = final_usage;
138                }
139                StreamEvent::Error { error } => {
140                    return Err(anyhow!("Stream error: {}", error.message));
141                }
142                _ => {}
143            }
144        }
145
146        info!(
147            "Streaming completed. Tokens: {} in, {} out",
148            usage.input_tokens, usage.output_tokens
149        );
150
151        Ok((text, usage))
152    }
153
154    /// Build request headers
155    fn build_headers(&self, config: &AnthropicConfig) -> Result<HeaderMap> {
156        let mut headers = HeaderMap::new();
157
158        headers.insert(
159            CONTENT_TYPE,
160            HeaderValue::from_static("application/json"),
161        );
162
163        headers.insert(
164            "x-api-key",
165            HeaderValue::from_str(&config.api_key)
166                .context("Invalid API key")?,
167        );
168
169        headers.insert(
170            "anthropic-version",
171            HeaderValue::from_str(&config.api_version)
172                .context("Invalid API version")?,
173        );
174
175        Ok(headers)
176    }
177
178    /// Parse Server-Sent Events stream
179    fn parse_sse_stream(
180        &self,
181        byte_stream: impl Stream<Item = reqwest::Result<bytes::Bytes>> + Send + 'static,
182    ) -> impl Stream<Item = Result<StreamEvent>> + Send {
183        let buffer = Arc::new(tokio::sync::Mutex::new(String::new()));
184
185        byte_stream.filter_map(move |chunk_result| {
186            let buffer = buffer.clone();
187            async move {
188                let chunk = match chunk_result {
189                    Ok(c) => c,
190                    Err(e) => return Some(Err(anyhow!("Stream error: {}", e))),
191                };
192
193                // Convert bytes to string
194                let chunk_str = match std::str::from_utf8(&chunk) {
195                    Ok(s) => s,
196                    Err(e) => return Some(Err(anyhow!("Invalid UTF-8: {}", e))),
197                };
198
199                let mut buffer = buffer.lock().await;
200                buffer.push_str(chunk_str);
201
202                // Process complete SSE messages
203                let mut events = Vec::new();
204
205                while let Some(pos) = buffer.find("\n\n") {
206                    let message = buffer[..pos].to_string();
207                    *buffer = buffer[pos + 2..].to_string();
208
209                    if message.is_empty() {
210                        continue;
211                    }
212
213                    // Parse SSE message
214                    let mut event_type = None;
215                    let mut data = String::new();
216
217                    for line in message.lines() {
218                        if let Some(stripped) = line.strip_prefix("event: ") {
219                            event_type = Some(stripped.to_string());
220                        } else if let Some(stripped) = line.strip_prefix("data: ") {
221                            data.push_str(stripped);
222                        }
223                    }
224
225                    // Parse event data as JSON
226                    if !data.is_empty() {
227                        match serde_json::from_str::<StreamEvent>(&data) {
228                            Ok(event) => {
229                                debug!("Received stream event: {:?}", event_type);
230                                events.push(Ok(event));
231                            }
232                            Err(e) => {
233                                warn!("Failed to parse stream event: {}", e);
234                                events.push(Err(anyhow!("Parse error: {}", e)));
235                            }
236                        }
237                    }
238                }
239
240                if events.is_empty() {
241                    None
242                } else if events.len() == 1 {
243                    Some(events.into_iter().next().unwrap())
244                } else {
245                    // If multiple events, return the first one
246                    // (This is a simplification; in practice, you might want to handle this differently)
247                    Some(events.into_iter().next().unwrap())
248                }
249            }
250        })
251    }
252}
253
254/// Stream collector for aggregating streaming events
255pub struct StreamCollector {
256    /// Accumulated text content
257    pub text: String,
258    /// Message ID
259    pub message_id: Option<String>,
260    /// Model used
261    pub model: Option<String>,
262    /// Usage statistics
263    pub usage: Usage,
264    /// Stop reason
265    pub stop_reason: Option<StopReason>,
266    /// Stop sequence
267    pub stop_sequence: Option<String>,
268}
269
270impl StreamCollector {
271    /// Create a new stream collector
272    pub fn new() -> Self {
273        Self {
274            text: String::new(),
275            message_id: None,
276            model: None,
277            usage: Usage {
278                input_tokens: 0,
279                output_tokens: 0,
280            },
281            stop_reason: None,
282            stop_sequence: None,
283        }
284    }
285
286    /// Process a stream event
287    ///
288    /// # Arguments
289    ///
290    /// * `event` - Stream event to process
291    ///
292    /// # Returns
293    ///
294    /// Returns true if this is the final event
295    pub fn process_event(&mut self, event: StreamEvent) -> bool {
296        match event {
297            StreamEvent::MessageStart { message } => {
298                self.message_id = Some(message.id);
299                self.model = Some(message.model);
300                self.usage = message.usage;
301                false
302            }
303            StreamEvent::ContentBlockDelta { delta, .. } => {
304                if let Delta::TextDelta { text } = delta {
305                    self.text.push_str(&text);
306                }
307                false
308            }
309            StreamEvent::MessageDelta { delta, usage } => {
310                self.usage = usage;
311                self.stop_reason = delta.stop_reason;
312                self.stop_sequence = delta.stop_sequence;
313                false
314            }
315            StreamEvent::MessageStop => {
316                true // Final event
317            }
318            StreamEvent::Error { error } => {
319                warn!("Stream error: {}", error.message);
320                true
321            }
322            _ => false,
323        }
324    }
325
326    /// Convert to MessageResponse
327    pub fn to_response(self) -> Result<MessageResponse> {
328        Ok(MessageResponse {
329            id: self.message_id.ok_or_else(|| anyhow!("Missing message ID"))?,
330            type_field: "message".to_string(),
331            role: Role::Assistant,
332            content: vec![ContentBlock::Text { text: self.text }],
333            model: self.model.ok_or_else(|| anyhow!("Missing model"))?,
334            stop_reason: self.stop_reason,
335            stop_sequence: self.stop_sequence,
336            usage: self.usage,
337        })
338    }
339}
340
341impl Default for StreamCollector {
342    fn default() -> Self {
343        Self::new()
344    }
345}
346
347#[cfg(test)]
348mod tests {
349    use super::*;
350
351    #[test]
352    fn test_stream_collector() {
353        let mut collector = StreamCollector::new();
354
355        // Process message start
356        let start_event = StreamEvent::MessageStart {
357            message: MessageStart {
358                id: "msg_123".to_string(),
359                type_field: "message".to_string(),
360                role: Role::Assistant,
361                content: vec![],
362                model: "claude-3-haiku-20240307".to_string(),
363                usage: Usage {
364                    input_tokens: 10,
365                    output_tokens: 0,
366                },
367            },
368        };
369
370        assert!(!collector.process_event(start_event));
371        assert_eq!(collector.message_id, Some("msg_123".to_string()));
372
373        // Process content delta
374        let delta_event = StreamEvent::ContentBlockDelta {
375            index: 0,
376            delta: Delta::TextDelta {
377                text: "Hello".to_string(),
378            },
379        };
380
381        assert!(!collector.process_event(delta_event));
382        assert_eq!(collector.text, "Hello");
383
384        // Process another delta
385        let delta_event2 = StreamEvent::ContentBlockDelta {
386            index: 0,
387            delta: Delta::TextDelta {
388                text: " world".to_string(),
389            },
390        };
391
392        assert!(!collector.process_event(delta_event2));
393        assert_eq!(collector.text, "Hello world");
394
395        // Process message stop
396        let stop_event = StreamEvent::MessageStop;
397        assert!(collector.process_event(stop_event));
398    }
399
400    #[test]
401    fn test_stream_collector_to_response() {
402        let mut collector = StreamCollector::new();
403
404        collector.message_id = Some("msg_123".to_string());
405        collector.model = Some("claude-3-haiku-20240307".to_string());
406        collector.text = "Hello, world!".to_string();
407        collector.usage = Usage {
408            input_tokens: 10,
409            output_tokens: 5,
410        };
411        collector.stop_reason = Some(StopReason::EndTurn);
412
413        let response = collector.to_response().unwrap();
414
415        assert_eq!(response.id, "msg_123");
416        assert_eq!(response.usage.input_tokens, 10);
417        assert_eq!(response.usage.output_tokens, 5);
418    }
419}