llm_optimizer_integrations/anthropic/
streaming.rs1use 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
14pub struct StreamHandler {
16 config: Arc<RwLock<AnthropicConfig>>,
18 client: reqwest::Client,
20 cost_tracker: Arc<RwLock<CostTracker>>,
22}
23
24impl StreamHandler {
25 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 pub async fn stream_message(
54 &self,
55 mut request: MessageRequest,
56 ) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent>> + Send>>> {
57 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 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 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 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 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 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 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 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 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 Some(events.into_iter().next().unwrap())
248 }
249 }
250 })
251 }
252}
253
254pub struct StreamCollector {
256 pub text: String,
258 pub message_id: Option<String>,
260 pub model: Option<String>,
262 pub usage: Usage,
264 pub stop_reason: Option<StopReason>,
266 pub stop_sequence: Option<String>,
268}
269
270impl StreamCollector {
271 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 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 }
318 StreamEvent::Error { error } => {
319 warn!("Stream error: {}", error.message);
320 true
321 }
322 _ => false,
323 }
324 }
325
326 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 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 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 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 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}