1use super::super::models::{Content, Part};
7use super::{StreamingCandidate, StreamingError, StreamingMetrics, StreamingResponse};
8use futures::stream::StreamExt;
9use reqwest::Response;
10use serde_json::Value;
11use std::time::Instant;
12use tokio::time::{Duration, timeout};
13use tracing;
14
15#[derive(Debug, Clone)]
17pub struct StreamingConfig {
18 chunk_timeout: Duration,
20 first_chunk_timeout: Duration,
22 total_timeout: Duration,
24 buffer_size: usize,
26}
27
28impl Default for StreamingConfig {
29 fn default() -> Self {
30 Self {
31 chunk_timeout: Duration::from_secs(30),
32 first_chunk_timeout: Duration::from_secs(60),
33 total_timeout: Duration::from_secs(600),
34 buffer_size: 1024,
35 }
36 }
37}
38
39impl StreamingConfig {
40 pub(crate) fn with_total_timeout(total_timeout_secs: u64) -> Self {
42 Self {
43 total_timeout: Duration::from_secs(total_timeout_secs),
44 ..Default::default()
45 }
46 }
47}
48
49pub type ProgressCallback = Box<dyn Fn(f32) + Send + Sync>;
51
52pub struct StreamingProcessor {
54 config: StreamingConfig,
55 metrics: StreamingMetrics,
56 current_event_data: String,
57 progress_callback: Option<ProgressCallback>,
58 warning_threshold: f32,
59}
60
61impl StreamingProcessor {
62 pub(crate) fn new() -> Self {
64 Self {
65 config: StreamingConfig::default(),
66 metrics: StreamingMetrics::default(),
67 current_event_data: String::with_capacity(512),
69 progress_callback: None,
70 warning_threshold: 0.8,
71 }
72 }
73
74 pub(crate) fn with_config(config: StreamingConfig) -> Self {
76 Self {
77 config,
78 metrics: StreamingMetrics::default(),
79 current_event_data: String::with_capacity(512),
81 progress_callback: None,
82 warning_threshold: 0.8,
83 }
84 }
85
86 pub fn with_progress_callback(mut self, callback: ProgressCallback) -> Self {
88 self.progress_callback = Some(callback);
89 self
90 }
91
92 pub fn with_warning_threshold(mut self, threshold: f32) -> Self {
94 self.warning_threshold = threshold.clamp(0.0, 1.0);
95 self
96 }
97
98 pub(crate) async fn process_stream<F>(
112 &mut self,
113 response: Response,
114 mut on_chunk: F,
115 ) -> Result<StreamingResponse, StreamingError>
116 where
117 F: FnMut(&str) -> Result<(), StreamingError>,
118 {
119 self.metrics.request_start_time = Some(Instant::now());
120 self.metrics.total_requests += 1;
121 self.current_event_data.clear();
122
123 let mut stream = response.bytes_stream();
125
126 let mut accumulated_response = StreamingResponse { candidates: Vec::new(), usage_metadata: None };
127
128 let mut _has_valid_content = false;
129 let mut buffer = String::with_capacity(self.config.buffer_size);
131 let mut decoder = crate::providers::shared::Utf8StreamDecoder::new();
132 let request_start = Instant::now();
133
134 let first_chunk_start = Instant::now();
136 let first_chunk_result = timeout(self.config.first_chunk_timeout, stream.next()).await;
137
138 match first_chunk_result {
139 Ok(Some(Ok(bytes))) => {
140 self.metrics.first_chunk_time = Some(Instant::now());
141 self.metrics.total_bytes += bytes.len();
142 self.report_progress(first_chunk_start, request_start);
143
144 buffer.push_str(&decoder.push(&bytes));
146 match self.process_buffer(&mut buffer, &mut accumulated_response, &mut on_chunk) {
147 Ok(valid) => _has_valid_content = valid,
148 Err(e) => return Err(e),
149 }
150 }
151 Ok(Some(Err(e))) => {
152 self.metrics.error_count += 1;
153 return Err(StreamingError::NetworkError {
154 message: format!("Failed to read first chunk: {e}"),
155 is_retryable: true,
156 });
157 }
158 Ok(None) => {
159 return Err(StreamingError::StreamingError {
160 message: "Empty streaming response".to_owned(),
161 partial_content: None,
162 });
163 }
164 Err(_) => {
165 self.metrics.error_count += 1;
166 self.report_progress(first_chunk_start, request_start);
167 return Err(StreamingError::TimeoutError {
168 operation: "first_chunk".to_owned(),
169 duration: self.config.first_chunk_timeout,
170 });
171 }
172 }
173
174 while let Some(result) = stream.next().await {
176 let elapsed = request_start.elapsed();
177
178 if self.config.total_timeout.as_secs() > 0 {
180 if elapsed > self.config.total_timeout {
181 self.metrics.error_count += 1;
182 self.report_progress_at_timeout(elapsed);
183 return Err(StreamingError::TimeoutError {
184 operation: "streaming".to_owned(),
185 duration: elapsed,
186 });
187 }
188 self.report_progress_with_timeout(elapsed);
190 }
191
192 match result {
193 Ok(bytes) => {
194 self.metrics.total_bytes += bytes.len();
195
196 buffer.push_str(&decoder.push(&bytes));
198
199 match self.process_buffer(&mut buffer, &mut accumulated_response, &mut on_chunk) {
201 Ok(valid) => {
202 if valid {
203 _has_valid_content = true;
204 }
205 }
206 Err(e) => return Err(e),
207 }
208 }
209 Err(e) => {
210 self.metrics.error_count += 1;
211 self.report_progress_at_timeout(elapsed);
212 return Err(StreamingError::NetworkError {
213 message: format!("Failed to read chunk: {e}"),
214 is_retryable: true,
215 });
216 }
217 }
218
219 self.metrics.total_chunks += 1;
220 }
221
222 if !buffer.is_empty() {
224 match self.process_remaining_buffer(&mut buffer, &mut accumulated_response, &mut on_chunk) {
225 Ok(valid) => {
226 if valid {
227 _has_valid_content = true;
228 }
229 }
230 Err(e) => return Err(e),
231 }
232 }
233
234 if !_has_valid_content {
235 return Err(StreamingError::ContentError {
236 message: "No valid content received from streaming API".to_owned(),
237 });
238 }
239
240 Ok(accumulated_response)
241 }
242
243 fn process_buffer<F>(
244 &mut self,
245 buffer: &mut String,
246 accumulated_response: &mut StreamingResponse,
247 on_chunk: &mut F,
248 ) -> Result<bool, StreamingError>
249 where
250 F: FnMut(&str) -> Result<(), StreamingError>,
251 {
252 let mut _has_valid_content = false;
253 let mut processed_chars = 0;
254
255 while let Some(newline_pos) = buffer[processed_chars..].find('\n') {
256 let line_end = processed_chars + newline_pos;
257 let line = &buffer[processed_chars..line_end];
258 processed_chars = line_end + 1;
259
260 _has_valid_content |= self.handle_line(line, accumulated_response, on_chunk)?;
261 }
262
263 if processed_chars > 0 {
265 buffer.drain(..processed_chars);
266 }
267
268 Ok(_has_valid_content)
269 }
270
271 fn process_remaining_buffer<F>(
273 &mut self,
274 buffer: &mut String,
275 accumulated_response: &mut StreamingResponse,
276 on_chunk: &mut F,
277 ) -> Result<bool, StreamingError>
278 where
279 F: FnMut(&str) -> Result<(), StreamingError>,
280 {
281 let mut _has_valid_content = false;
282
283 if !buffer.is_empty() {
284 let remaining_line = buffer.trim_end_matches('\r');
285 if !remaining_line.trim().is_empty() {
286 _has_valid_content |= self.handle_line(remaining_line, accumulated_response, on_chunk)?;
287 }
288 }
289
290 buffer.clear();
291
292 _has_valid_content |= self.finalize_current_event(accumulated_response, on_chunk)?;
293
294 Ok(_has_valid_content)
295 }
296
297 fn handle_line<F>(
299 &mut self,
300 raw_line: &str,
301 accumulated_response: &mut StreamingResponse,
302 on_chunk: &mut F,
303 ) -> Result<bool, StreamingError>
304 where
305 F: FnMut(&str) -> Result<(), StreamingError>,
306 {
307 let mut _has_valid_content = false;
308 let line = raw_line.trim_end_matches('\r');
309
310 if line.is_empty() {
311 _has_valid_content |= self.finalize_current_event(accumulated_response, on_chunk)?;
312 return Ok(_has_valid_content);
313 }
314
315 let trimmed = line.trim();
316
317 if trimmed.is_empty() {
318 return Ok(false);
319 }
320
321 if trimmed.starts_with(':') {
322 return Ok(false);
323 }
324
325 if trimmed.starts_with("event:") || trimmed.starts_with("id:") {
326 return Ok(false);
327 }
328
329 if let Some(data_segment) = trimmed.strip_prefix("data:") {
330 let data_segment = data_segment.trim_start();
331 if data_segment == "[DONE]" {
332 _has_valid_content |= self.finalize_current_event(accumulated_response, on_chunk)?;
333 return Ok(_has_valid_content);
334 }
335
336 if !data_segment.is_empty() {
337 if !self.current_event_data.is_empty() {
338 self.current_event_data.push('\n');
339 }
340 self.current_event_data.push_str(data_segment);
341
342 _has_valid_content |= self.try_flush_current_event(accumulated_response, on_chunk)?;
343 }
344 return Ok(_has_valid_content);
345 }
346
347 if trimmed.starts_with('{') || trimmed.starts_with('[') {
348 if !self.current_event_data.is_empty() {
349 self.current_event_data.push('\n');
350 }
351 self.current_event_data.push_str(trimmed);
352 return Ok(false);
353 }
354
355 if !self.current_event_data.is_empty() {
356 self.current_event_data.push('\n');
357 }
358 self.current_event_data.push_str(trimmed);
359
360 Ok(false)
361 }
362
363 fn finalize_current_event<F>(
364 &mut self,
365 accumulated_response: &mut StreamingResponse,
366 on_chunk: &mut F,
367 ) -> Result<bool, StreamingError>
368 where
369 F: FnMut(&str) -> Result<(), StreamingError>,
370 {
371 if self.current_event_data.trim().is_empty() {
372 self.current_event_data.clear();
373 return Ok(false);
374 }
375
376 let event_data = std::mem::take(&mut self.current_event_data);
377 self.process_event(event_data, accumulated_response, on_chunk)
378 }
379
380 fn try_flush_current_event<F>(
381 &mut self,
382 accumulated_response: &mut StreamingResponse,
383 on_chunk: &mut F,
384 ) -> Result<bool, StreamingError>
385 where
386 F: FnMut(&str) -> Result<(), StreamingError>,
387 {
388 let trimmed = self.current_event_data.trim();
389 if trimmed.is_empty() {
390 return Ok(false);
391 }
392
393 match serde_json::from_str::<Value>(trimmed) {
394 Ok(parsed) => {
395 self.current_event_data.clear();
396 self.process_event_value(parsed, accumulated_response, on_chunk)
397 }
398 Err(parse_err) => {
399 if parse_err.is_eof() {
400 return Ok(false);
401 }
402
403 Err(StreamingError::ParseError {
404 message: format!("Failed to parse streaming JSON: {parse_err}"),
405 raw_response: trimmed.to_owned(),
406 })
407 }
408 }
409 }
410
411 fn process_event<F>(
412 &mut self,
413 event_data: String,
414 accumulated_response: &mut StreamingResponse,
415 on_chunk: &mut F,
416 ) -> Result<bool, StreamingError>
417 where
418 F: FnMut(&str) -> Result<(), StreamingError>,
419 {
420 let trimmed = event_data.trim();
421
422 if trimmed.is_empty() {
423 return Ok(false);
424 }
425
426 match serde_json::from_str::<Value>(trimmed) {
427 Ok(parsed) => self.process_event_value(parsed, accumulated_response, on_chunk),
428 Err(parse_err) => {
429 if parse_err.is_eof() {
430 self.current_event_data = trimmed.to_owned();
431 return Ok(false);
432 }
433
434 Err(StreamingError::ParseError {
435 message: format!("Failed to parse streaming JSON: {parse_err}"),
436 raw_response: trimmed.to_owned(),
437 })
438 }
439 }
440 }
441
442 fn append_text_candidate(&mut self, accumulated_response: &mut StreamingResponse, text: &str) {
443 if text.is_empty() {
444 return;
445 }
446
447 if let Some(last_candidate) = accumulated_response.candidates.last_mut() {
448 Self::merge_parts(
449 &mut last_candidate.content.parts,
450 vec![Part::Text { text: text.to_owned(), thought_signature: None }],
451 );
452 return;
453 }
454
455 let index = accumulated_response.candidates.len();
456
457 accumulated_response.candidates.push(StreamingCandidate {
458 content: Content {
459 role: "model".to_owned(),
460 parts: vec![Part::Text { text: text.to_owned(), thought_signature: None }],
461 },
462 finish_reason: None,
463 index: Some(index),
464 });
465 }
466
467 fn process_candidate<F>(&self, candidate: &StreamingCandidate, on_chunk: &mut F) -> Result<bool, StreamingError>
469 where
470 F: FnMut(&str) -> Result<(), StreamingError>,
471 {
472 let mut _has_valid_content = false;
473
474 if candidate.finish_reason.is_some() {
475 _has_valid_content = true;
476 }
477
478 for part in &candidate.content.parts {
480 match part {
481 Part::Text { text, .. } => {
482 if !text.trim().is_empty() {
483 on_chunk(text)?;
484 _has_valid_content = true;
485 }
486 }
487 Part::InlineData { .. } => {
488 _has_valid_content = true;
489 }
490 Part::FunctionCall { .. } => {
491 _has_valid_content = true;
493 }
494 Part::FunctionResponse { .. } => {
495 _has_valid_content = true;
496 }
497 Part::ToolCall { .. }
498 | Part::ToolResponse { .. }
499 | Part::ExecutableCode { .. }
500 | Part::CodeExecutionResult { .. } => {
501 _has_valid_content = true;
502 }
503 Part::CacheControl { .. } => {}
504 }
505 }
506
507 Ok(_has_valid_content)
508 }
509
510 fn process_event_value<F>(
511 &mut self,
512 value: Value,
513 accumulated_response: &mut StreamingResponse,
514 on_chunk: &mut F,
515 ) -> Result<bool, StreamingError>
516 where
517 F: FnMut(&str) -> Result<(), StreamingError>,
518 {
519 match value {
520 Value::Array(items) => {
521 let mut has_valid = false;
522 for item in items {
523 if self.process_event_value(item, accumulated_response, on_chunk)? {
524 has_valid = true;
525 }
526 }
527 Ok(has_valid)
528 }
529 Value::Object(mut map) => {
530 if let Some(error_value) = map.get("error") {
531 let message = error_value
532 .get("message")
533 .and_then(Value::as_str)
534 .unwrap_or("Gemini streaming error")
535 .to_owned();
536 #[allow(
537 clippy::cast_sign_loss,
538 reason = "Intentional compatibility, platform, or test-only suppression."
539 )]
540 let code = error_value.get("code").and_then(Value::as_i64).unwrap_or(500) as u16;
541 return Err(StreamingError::ApiError {
542 status_code: code,
543 message,
544 is_retryable: code == 429,
545 });
546 }
547
548 if let Some(usage) = map.remove("usageMetadata") {
549 accumulated_response.usage_metadata = Some(usage);
550 }
551
552 let mut has_valid = false;
553
554 if let Some(candidates_value) = map.remove("candidates") {
555 let candidate_values: Vec<Value> = match candidates_value {
556 Value::Array(items) => items,
557 Value::Object(_) => vec![candidates_value],
558 _ => Vec::new(),
559 };
560
561 for candidate_value in candidate_values {
562 match serde_json::from_value::<StreamingCandidate>(candidate_value.clone()) {
563 Ok(candidate) => {
564 if self.process_candidate(&candidate, on_chunk)? {
565 has_valid = true;
566 }
567 self.merge_candidate(accumulated_response, candidate);
568 }
569 Err(err) => {
570 if let Some(text) = Self::extract_text_from_value(&candidate_value) {
571 if !text.trim().is_empty() {
572 on_chunk(&text)?;
573 self.append_text_candidate(accumulated_response, &text);
574 has_valid = true;
575 }
576 } else {
577 return Err(StreamingError::ParseError {
578 message: format!("Failed to parse candidate: {err}"),
579 raw_response: candidate_value.to_string(),
580 });
581 }
582 }
583 }
584 }
585 }
586
587 if let Some(text_value) = map.remove("text").and_then(|v| v.as_str().map(|s| s.to_owned()))
588 && !text_value.trim().is_empty()
589 {
590 on_chunk(&text_value)?;
591 self.append_text_candidate(accumulated_response, &text_value);
592 has_valid = true;
593 }
594
595 Ok(has_valid)
596 }
597 Value::String(text) => {
598 if text.trim().is_empty() {
599 Ok(false)
600 } else {
601 on_chunk(&text)?;
602 self.append_text_candidate(accumulated_response, &text);
603 Ok(true)
604 }
605 }
606 _ => Ok(false),
607 }
608 }
609
610 fn merge_candidate(&mut self, accumulated_response: &mut StreamingResponse, mut candidate: StreamingCandidate) {
611 let index = candidate.index.unwrap_or(accumulated_response.candidates.len());
612
613 if let Some(existing) = accumulated_response
614 .candidates
615 .iter_mut()
616 .find(|existing| existing.index.unwrap_or(index) == index)
617 {
618 if existing.content.role.is_empty() {
619 existing.content.role = candidate.content.role.clone();
620 }
621
622 Self::merge_parts(&mut existing.content.parts, candidate.content.parts);
623
624 if candidate.finish_reason.is_some() {
625 existing.finish_reason = candidate.finish_reason;
626 }
627 } else {
628 candidate.index = Some(index);
629 accumulated_response.candidates.push(candidate);
630 }
631 }
632
633 fn merge_parts(target: &mut Vec<Part>, source_parts: Vec<Part>) {
634 if target.is_empty() {
635 *target = source_parts;
636 return;
637 }
638
639 for part in source_parts {
640 match (target.last_mut(), &part) {
641 (
642 Some(Part::Text { text: existing, thought_signature: existing_sig }),
643 Part::Text { text: new_text, thought_signature: new_sig },
644 ) => {
645 existing.push_str(new_text);
646 if existing_sig.is_none() && new_sig.is_some() {
648 *existing_sig = new_sig.clone();
649 }
650 }
651 _ => target.push(part),
652 }
653 }
654 }
655
656 fn extract_text_from_value(value: &Value) -> Option<String> {
657 match value {
658 Value::String(text) => {
659 if text.trim().is_empty() {
660 None
661 } else {
662 Some(text.clone())
663 }
664 }
665 Value::Array(items) => Self::extract_text_from_array(items),
666 Value::Object(map) => {
667 if let Some(text) = map.get("text").and_then(Value::as_str)
668 && !text.trim().is_empty()
669 {
670 return Some(text.to_owned());
671 }
672
673 if let Some(parts) = map.get("parts").and_then(Value::as_array)
674 && let Some(parts_text) = Self::extract_text_from_array(parts)
675 {
676 return Some(parts_text);
677 }
678
679 for nested in map.values() {
680 if let Some(nested_text) = Self::extract_text_from_value(nested)
681 && !nested_text.trim().is_empty()
682 {
683 return Some(nested_text);
684 }
685 }
686
687 None
688 }
689 _ => None,
690 }
691 }
692
693 fn extract_text_from_array(items: &[Value]) -> Option<String> {
697 let mut collected = String::new();
698 for item in items {
699 if let Some(fragment) = Self::extract_text_from_value(item) {
700 collected.push_str(&fragment);
701 }
702 }
703 if collected.is_empty() { None } else { Some(collected) }
704 }
705
706 fn report_progress_with_timeout(&self, elapsed: Duration) {
708 if self.config.total_timeout.as_secs() == 0 {
709 return;
710 }
711
712 let progress = elapsed.as_secs_f32() / self.config.total_timeout.as_secs_f32();
713 let progress_clamped = progress.min(0.99); if let Some(ref callback) = self.progress_callback {
716 callback(progress_clamped);
717 }
718
719 if progress >= self.warning_threshold {
721 tracing::warn!(
722 "Streaming operation at {:.0}% of timeout limit ({}/{:?} elapsed). Approaching timeout.",
723 progress_clamped * 100.0,
724 elapsed.as_secs(),
725 self.config.total_timeout
726 );
727 }
728 }
729
730 fn report_progress_at_timeout(&self, _elapsed: Duration) {
732 if let Some(ref callback) = self.progress_callback {
733 callback(1.0); }
735 }
736
737 fn report_progress(&self, _event_time: Instant, _start_time: Instant) {
739 if let Some(ref callback) = self.progress_callback {
740 callback(0.1); }
742 }
743
744 fn metrics(&self) -> &StreamingMetrics {
746 &self.metrics
747 }
748
749 pub fn reset_metrics(&mut self) {
751 self.metrics = StreamingMetrics::default();
752 }
753}
754
755impl Default for StreamingProcessor {
756 fn default() -> Self {
757 Self::new()
758 }
759}
760
761#[cfg(test)]
762mod tests {
763 use super::*;
764
765 #[test]
766 fn test_streaming_processor_creation() {
767 let processor = StreamingProcessor::new();
768 assert_eq!(processor.metrics().total_requests, 0);
769 }
770
771 #[test]
772 fn test_streaming_processor_with_config() {
773 use std::time::Duration;
774
775 let config = StreamingConfig {
776 chunk_timeout: Duration::from_secs(10),
777 first_chunk_timeout: Duration::from_secs(30),
778 total_timeout: Duration::from_secs(120),
779 buffer_size: 512,
780 };
781
782 let processor = StreamingProcessor::with_config(config);
783 assert_eq!(processor.metrics().total_requests, 0);
784 }
785
786 #[test]
787 fn test_streaming_config_default() {
788 let config = StreamingConfig::default();
789 assert_eq!(config.buffer_size, 1024);
790 }
791
792 #[test]
793 fn test_handles_back_to_back_data_lines_without_blank_lines() {
794 let mut processor = StreamingProcessor::new();
795 let mut accumulated = StreamingResponse { candidates: Vec::new(), usage_metadata: None };
796 let mut received_chunks: Vec<String> = Vec::new();
797 let mut buffer = String::from(
798 "data: {\"candidates\":[{\"index\":0,\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"Hello\"}]}}]}\n",
799 );
800 buffer.push_str(
801 "data: {\"candidates\":[{\"index\":0,\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\" world\"}]}}]}\n",
802 );
803
804 {
805 let mut on_chunk = |chunk: &str| {
806 received_chunks.push(chunk.to_owned());
807 Ok(())
808 };
809 let has_valid = processor
810 .process_buffer(&mut buffer, &mut accumulated, &mut on_chunk)
811 .expect("processing should succeed");
812 assert!(has_valid);
813 }
814
815 assert_eq!(received_chunks, vec!["Hello", " world"]);
816 assert_eq!(accumulated.candidates.len(), 1);
817 let combined = match &accumulated.candidates[0].content.parts[0] {
818 Part::Text { text, .. } => text.clone(),
819 _ => String::new(),
820 };
821 assert_eq!(combined, "Hello world");
822 }
823}