1use std::collections::VecDeque;
13use std::sync::{Arc, Mutex};
14
15use futures_util::stream::BoxStream;
16use hotl_types::{Item, StopReason, TokenUsage};
17use serde::{Deserialize, Serialize};
18use serde_json::Value;
19
20pub mod timeouts {
28 use std::time::Duration;
29
30 pub const CONNECT: Duration = Duration::from_secs(10);
33
34 pub const HEADERS: Duration = Duration::from_secs(120);
37
38 pub const STREAM_IDLE: Duration = Duration::from_secs(300);
42}
43
44#[derive(Debug, Clone, Serialize, Deserialize)]
45pub struct ToolDef {
46 pub name: String,
47 pub description: String,
48 pub input_schema: Value,
49}
50
51#[derive(Debug, Clone)]
52pub struct SamplingRequest {
53 pub model: String,
54 pub max_tokens: u32,
55 pub system: Arc<str>,
57 pub items: Arc<Vec<Item>>,
61 pub ephemeral_tail: Arc<Vec<Item>>,
67 pub tools: Arc<[ToolDef]>,
68 pub thinking: bool,
70 pub cache: CachePolicy,
72 pub turn_context: Option<String>,
76}
77
78#[derive(Debug, Clone, Copy, PartialEq, Eq)]
81pub enum CacheTtl {
82 FiveMinutes,
83 OneHour,
84}
85
86#[derive(Debug, Clone, Copy, PartialEq, Eq)]
92pub enum CachePolicy {
93 Off,
96 Static { prefix_ttl: CacheTtl },
103}
104
105impl CachePolicy {
106 pub fn marks_breakpoints(self) -> bool {
110 matches!(self, Self::Static { .. })
111 }
112}
113
114#[derive(Debug, Clone, Serialize, Deserialize)]
116#[serde(tag = "event", rename_all = "snake_case")]
117pub enum StreamEvent {
118 Started,
119 BlockStart {
120 index: usize,
121 kind: String,
122 },
123 TextDelta {
124 index: usize,
125 text: String,
126 },
127 ThinkingDelta {
128 index: usize,
129 text: String,
130 },
131 ToolInputDelta {
132 index: usize,
133 json: String,
134 },
135 BlockEnd {
136 index: usize,
137 },
138 Retrying {
139 attempt: u32,
140 reason: String,
141 },
142 Completed {
145 stop: StopReason,
146 usage: TokenUsage,
147 blocks: Vec<Value>,
148 },
149}
150
151#[derive(Debug, Clone, thiserror::Error)]
154pub enum ProviderError {
155 #[error("authentication failed: {0}")]
156 Auth(String),
157 #[error("HTTP {status}: {message}")]
158 Http {
159 status: u16,
160 message: String,
161 retry_after: Option<u64>,
162 },
163 #[error("transport error: {0}")]
164 Transport(String),
165 #[error("stream parse error: {0}")]
166 Parse(String),
167}
168
169pub trait Provider: Send + Sync {
170 fn stream(
171 &self,
172 req: SamplingRequest,
173 ) -> BoxStream<'static, Result<StreamEvent, ProviderError>>;
174
175 fn arm(&self) -> ArmGuard {
180 ArmGuard::noop()
181 }
182}
183
184pub trait Warmable {
193 fn arm(&self) -> ArmGuard;
194}
195
196#[must_use]
210pub struct ArmGuard {
211 cancel: Option<Box<dyn FnOnce() + Send>>,
212}
213
214impl ArmGuard {
215 pub fn noop() -> Self {
217 Self { cancel: None }
218 }
219
220 pub fn new(cancel: impl FnOnce() + Send + 'static) -> Self {
222 Self {
223 cancel: Some(Box::new(cancel)),
224 }
225 }
226
227 pub fn detach(mut self) {
232 self.cancel = None;
233 }
234}
235
236impl Drop for ArmGuard {
237 fn drop(&mut self) {
238 if let Some(cancel) = self.cancel.take() {
239 cancel();
240 }
241 }
242}
243
244#[cfg(test)]
245mod arm_guard_tests {
246 use super::*;
247 use std::sync::atomic::{AtomicBool, Ordering};
248
249 #[test]
251 fn drop_invokes_the_cancel_callback() {
252 let called = Arc::new(AtomicBool::new(false));
253 let flag = called.clone();
254 let guard = ArmGuard::new(move || flag.store(true, Ordering::SeqCst));
255 assert!(!called.load(Ordering::SeqCst));
256 drop(guard);
257 assert!(called.load(Ordering::SeqCst));
258 }
259
260 #[test]
263 fn noop_cancels_nothing() {
264 drop(ArmGuard::noop());
265 }
266
267 #[test]
271 fn detach_suppresses_the_cancel_callback() {
272 let called = Arc::new(AtomicBool::new(false));
273 let flag = called.clone();
274 let guard = ArmGuard::new(move || flag.store(true, Ordering::SeqCst));
275 guard.detach();
276 assert!(!called.load(Ordering::SeqCst));
277 }
278
279 #[test]
284 fn provider_default_arm_is_a_noop() {
285 let provider = ScriptedProvider::new(vec![]);
286 drop(provider.arm());
288 }
289}
290
291pub struct ScriptedProvider {
294 scripts: ScriptQueue,
297 requests: Mutex<Vec<SamplingRequest>>,
299}
300
301type Script = Vec<Result<StreamEvent, ProviderError>>;
303type ScriptQueue = Arc<Mutex<VecDeque<Script>>>;
306
307struct ScriptedStream {
318 script: Script,
319 pos: usize,
320 exhausted: bool,
322 scripts: ScriptQueue,
323}
324
325impl futures_util::Stream for ScriptedStream {
326 type Item = Result<StreamEvent, ProviderError>;
327
328 fn poll_next(
329 self: std::pin::Pin<&mut Self>,
330 _cx: &mut std::task::Context<'_>,
331 ) -> std::task::Poll<Option<Self::Item>> {
332 let this = self.get_mut();
333 match this.script.get(this.pos) {
334 Some(event) => {
335 this.pos += 1;
336 std::task::Poll::Ready(Some(event.clone()))
337 }
338 None => {
339 this.exhausted = true;
340 std::task::Poll::Ready(None)
341 }
342 }
343 }
344}
345
346impl Drop for ScriptedStream {
347 fn drop(&mut self) {
348 if self.exhausted {
349 return;
350 }
351 self.scripts
352 .lock()
353 .expect("scripted provider mutex")
354 .push_front(std::mem::take(&mut self.script));
355 }
356}
357
358impl ScriptedProvider {
359 pub fn new(scripts: Vec<Vec<Result<StreamEvent, ProviderError>>>) -> Self {
360 Self {
361 scripts: Arc::new(Mutex::new(scripts.into())),
362 requests: Mutex::new(Vec::new()),
363 }
364 }
365
366 pub fn requests(&self) -> Vec<SamplingRequest> {
369 self.requests.lock().expect("requests mutex").clone()
370 }
371
372 pub fn last_request(&self) -> Option<SamplingRequest> {
374 self.requests
375 .lock()
376 .expect("requests mutex")
377 .last()
378 .cloned()
379 }
380
381 pub fn request_count(&self) -> usize {
382 self.requests.lock().expect("requests mutex").len()
383 }
384
385 pub fn push_script(&self, script: Vec<Result<StreamEvent, ProviderError>>) {
388 self.scripts
389 .lock()
390 .expect("scripted provider mutex")
391 .push_back(script);
392 }
393
394 pub fn text_reply(text: &str) -> Vec<Result<StreamEvent, ProviderError>> {
396 vec![
397 Ok(StreamEvent::Started),
398 Ok(StreamEvent::BlockStart {
399 index: 0,
400 kind: "text".into(),
401 }),
402 Ok(StreamEvent::TextDelta {
403 index: 0,
404 text: text.into(),
405 }),
406 Ok(StreamEvent::BlockEnd { index: 0 }),
407 Ok(StreamEvent::Completed {
408 stop: StopReason::EndTurn,
409 usage: TokenUsage {
410 input_tokens: 10,
411 output_tokens: 5,
412 ..Default::default()
413 },
414 blocks: vec![serde_json::json!({"type": "text", "text": text})],
415 }),
416 ]
417 }
418
419 pub fn tool_call(
421 id: &str,
422 name: &str,
423 input: Value,
424 ) -> Vec<Result<StreamEvent, ProviderError>> {
425 let block = serde_json::json!({"type": "tool_use", "id": id, "name": name, "input": input});
426 vec![
427 Ok(StreamEvent::Started),
428 Ok(StreamEvent::BlockStart {
429 index: 0,
430 kind: "tool_use".into(),
431 }),
432 Ok(StreamEvent::BlockEnd { index: 0 }),
433 Ok(StreamEvent::Completed {
434 stop: StopReason::ToolUse,
435 usage: TokenUsage {
436 input_tokens: 10,
437 output_tokens: 8,
438 ..Default::default()
439 },
440 blocks: vec![block],
441 }),
442 ]
443 }
444}
445
446impl Provider for ScriptedProvider {
447 fn stream(
448 &self,
449 req: SamplingRequest,
450 ) -> BoxStream<'static, Result<StreamEvent, ProviderError>> {
451 self.requests.lock().expect("requests mutex").push(req);
452 let script = self
453 .scripts
454 .lock()
455 .expect("scripted provider mutex")
456 .pop_front()
457 .unwrap_or_else(|| {
458 vec![Err(ProviderError::Transport(
459 "scripted provider exhausted".into(),
460 ))]
461 });
462 Box::pin(ScriptedStream {
463 script,
464 pos: 0,
465 exhausted: false,
466 scripts: Arc::clone(&self.scripts),
467 })
468 }
469}
470
471pub fn v1_base(base: &str) -> String {
480 let base = base.trim_end_matches('/');
481 if base.ends_with("/v1") {
482 base.to_string()
483 } else {
484 format!("{base}/v1")
485 }
486}
487
488#[cfg(test)]
489mod base_url_tests {
490 use super::v1_base;
491
492 #[test]
493 fn both_spellings_and_trailing_slashes_resolve_alike() {
494 for input in [
495 "http://127.0.0.1:3456",
496 "http://127.0.0.1:3456/",
497 "http://127.0.0.1:3456/v1",
498 "http://127.0.0.1:3456/v1/",
499 ] {
500 assert_eq!(v1_base(input), "http://127.0.0.1:3456/v1", "input: {input}");
501 }
502 }
503}
504
505#[derive(Default)]
515pub struct SseParser {
516 buf: Vec<u8>,
517 data: Vec<String>,
519 data_len: usize,
520}
521
522pub const SSE_MAX_BUFFER: usize = 1024 * 1024;
524pub const SSE_MAX_EVENT: usize = 4 * 1024 * 1024;
526
527impl SseParser {
528 pub fn feed(&mut self, chunk: &[u8]) -> Result<Vec<String>, ProviderError> {
531 self.buf.extend_from_slice(chunk);
532 let mut out = Vec::new();
533 let mut start = 0;
534 while let Some(pos) = self.buf[start..].iter().position(|&b| b == b'\n') {
535 let line = String::from_utf8_lossy(&self.buf[start..start + pos])
536 .trim_end_matches('\r')
537 .to_string();
538 start += pos + 1;
539 self.line(&line, &mut out)?;
540 }
541 self.buf.drain(..start);
542 if self.buf.len() > SSE_MAX_BUFFER {
543 return Err(ProviderError::Parse(format!(
544 "SSE line exceeded {SSE_MAX_BUFFER} bytes without a newline"
545 )));
546 }
547 Ok(out)
548 }
549
550 pub fn finish(&mut self) -> Result<Vec<String>, ProviderError> {
555 let mut out = Vec::new();
556 if !self.buf.is_empty() {
557 let tail = std::mem::take(&mut self.buf);
558 let line = String::from_utf8_lossy(&tail)
559 .trim_end_matches('\r')
560 .to_string();
561 self.line(&line, &mut out)?;
562 }
563 self.dispatch(&mut out);
564 Ok(out)
565 }
566
567 fn line(&mut self, line: &str, out: &mut Vec<String>) -> Result<(), ProviderError> {
568 if line.is_empty() {
569 self.dispatch(out);
570 return Ok(());
571 }
572 if line.starts_with(':') {
573 return Ok(()); }
575 let Some(value) = line.strip_prefix("data:") else {
578 return Ok(());
579 };
580 let value = value.strip_prefix(' ').unwrap_or(value);
582 self.data_len += value.len() + 1;
583 if self.data_len > SSE_MAX_EVENT {
584 return Err(ProviderError::Parse(format!(
585 "SSE event exceeded {SSE_MAX_EVENT} bytes without a blank line"
586 )));
587 }
588 self.data.push(value.to_string());
589 Ok(())
590 }
591
592 fn dispatch(&mut self, out: &mut Vec<String>) {
593 self.data_len = 0;
594 if self.data.is_empty() {
595 return;
596 }
597 let payload = std::mem::take(&mut self.data).join("\n");
598 if !payload.is_empty() && payload != "[DONE]" {
599 out.push(payload);
600 }
601 }
602}
603
604#[cfg(test)]
605mod sse_parser_tests {
606 use super::*;
607
608 #[test]
612 fn multi_line_data_fields_are_joined() {
613 let mut p = SseParser::default();
614 let out = p
615 .feed(b"event: x\ndata: {\"a\":\ndata: 1}\n\ndata: {\"b\":2}\n\n")
616 .unwrap();
617 assert_eq!(
618 out,
619 vec!["{\"a\":\n1}".to_string(), "{\"b\":2}".to_string()]
620 );
621 for payload in &out {
623 serde_json::from_str::<serde_json::Value>(payload).expect(payload);
624 }
625 }
626
627 #[test]
630 fn the_unterminated_final_line_is_flushed() {
631 let mut p = SseParser::default();
632 assert!(p
633 .feed(b"data: {\"type\":\"message_stop\"}")
634 .unwrap()
635 .is_empty());
636 assert_eq!(p.finish().unwrap(), vec!["{\"type\":\"message_stop\"}"]);
637 }
638
639 #[test]
642 fn a_trailing_event_without_a_blank_line_is_flushed() {
643 let mut p = SseParser::default();
644 assert!(p.feed(b"data: {\"n\":1}\n").unwrap().is_empty());
645 assert_eq!(p.finish().unwrap(), vec!["{\"n\":1}"]);
646 }
647
648 #[test]
651 fn comments_other_fields_and_done_are_filtered() {
652 let mut p = SseParser::default();
653 let out = p
654 .feed(b": keepalive\nevent: ping\nid: 7\nretry: 100\ndata: [DONE]\n\ndata: {}\n\n")
655 .unwrap();
656 assert_eq!(out, vec!["{}".to_string()]);
657 }
658
659 #[test]
662 fn over_cap_input_is_a_parse_error_not_an_oom() {
663 let mut p = SseParser::default();
664 let big = vec![b'x'; SSE_MAX_BUFFER + 1];
665 assert!(matches!(p.feed(&big), Err(ProviderError::Parse(_))));
666
667 let mut q = SseParser::default();
668 let line = format!("data: {}\n", "y".repeat(64 * 1024));
669 let err = loop {
670 if let Err(e) = q.feed(line.as_bytes()) {
671 break e;
672 }
673 };
674 assert!(matches!(err, ProviderError::Parse(_)), "{err:?}");
675 }
676
677 #[test]
680 fn chunk_boundaries_and_utf8_splits_survive() {
681 let wire = "data: {\"t\":\"héllo → wörld\"}\n\n".as_bytes();
682 let mut p = SseParser::default();
683 let mut out = Vec::new();
684 for c in wire.chunks(3) {
685 out.extend(p.feed(c).unwrap());
686 }
687 out.extend(p.finish().unwrap());
688 assert_eq!(out.len(), 1);
689 let v: serde_json::Value = serde_json::from_str(&out[0]).unwrap();
690 assert_eq!(v["t"], "héllo → wörld");
691 }
692}
693
694pub mod retry {
697 use super::ProviderError;
698 use std::time::Duration;
699
700 #[derive(Debug, Clone, Copy, PartialEq, Eq)]
701 pub enum Decision {
702 Retry { delay: Duration },
705 Fatal,
707 }
708
709 pub const MAX_ATTEMPTS: u32 = 5;
712
713 pub const RETRY_AFTER_CAP: Duration = Duration::from_secs(60);
717
718 pub fn classify(err: &ProviderError, attempt: u32) -> Decision {
721 if attempt >= MAX_ATTEMPTS {
722 return Decision::Fatal;
723 }
724 let backoff = Duration::from_secs(1u64 << (attempt - 1));
725 match err {
726 ProviderError::Http {
727 status,
728 retry_after,
729 ..
730 } if *status == 429 || *status >= 500 => Decision::Retry {
731 delay: retry_after
732 .map(Duration::from_secs)
733 .unwrap_or(backoff)
734 .min(RETRY_AFTER_CAP),
735 },
736 ProviderError::Transport(_) => Decision::Retry {
737 delay: backoff.min(RETRY_AFTER_CAP),
738 },
739 _ => Decision::Fatal,
740 }
741 }
742
743 pub fn with_jitter(base: Duration) -> Duration {
754 use std::hash::{BuildHasher, Hasher};
755 let nanos = std::time::SystemTime::now()
756 .duration_since(std::time::UNIX_EPOCH)
757 .map(|d| d.subsec_nanos())
758 .unwrap_or(0);
759 let mut h = std::collections::hash_map::RandomState::new().build_hasher();
760 h.write_u32(nanos);
761 h.write_u32(std::process::id());
762 let half = base / 2;
763 let span = base.saturating_sub(half).as_nanos() as u64;
764 if span == 0 {
765 return base;
766 }
767 half + Duration::from_nanos(h.finish() % (span + 1))
768 }
769
770 pub fn now_unix() -> u64 {
772 std::time::SystemTime::now()
773 .duration_since(std::time::UNIX_EPOCH)
774 .map(|d| d.as_secs())
775 .unwrap_or(0)
776 }
777
778 pub fn parse_retry_after(value: &str, now_unix: u64) -> Option<u64> {
786 let v = value.trim();
787 if let Ok(secs) = v.parse::<u64>() {
788 return Some(secs);
789 }
790 let rest = v.split_once(", ")?.1;
792 let mut parts = rest.split(' ');
793 let day: u32 = parts.next()?.parse().ok()?;
794 let month = match parts.next()? {
795 "Jan" => 1,
796 "Feb" => 2,
797 "Mar" => 3,
798 "Apr" => 4,
799 "May" => 5,
800 "Jun" => 6,
801 "Jul" => 7,
802 "Aug" => 8,
803 "Sep" => 9,
804 "Oct" => 10,
805 "Nov" => 11,
806 "Dec" => 12,
807 _ => return None,
808 };
809 let year: i64 = parts.next()?.parse().ok()?;
810 let mut hms = parts.next()?.split(':');
811 let h: u64 = hms.next()?.parse().ok()?;
812 let m: u64 = hms.next()?.parse().ok()?;
813 let s: u64 = hms.next()?.parse().ok()?;
814 if h > 23 || m > 59 || s > 60 || !(1..=31).contains(&day) {
815 return None;
816 }
817 let secs =
818 days_from_civil(year, month, day).checked_mul(86_400)? as u64 + h * 3600 + m * 60 + s;
819 Some(secs.saturating_sub(now_unix))
820 }
821
822 fn days_from_civil(y: i64, m: u32, d: u32) -> i64 {
826 let y = if m <= 2 { y - 1 } else { y };
827 let era = if y >= 0 { y } else { y - 399 } / 400;
828 let yoe = y - era * 400; let mp = if m > 2 { m - 3 } else { m + 9 } as i64; let doy = (153 * mp + 2) / 5 + d as i64 - 1;
831 let doe = yoe * 365 + yoe / 4 - yoe / 100 + doy;
832 era * 146_097 + doe - 719_468
833 }
834
835 pub fn is_availability(err: &ProviderError) -> bool {
838 matches!(
839 err,
840 ProviderError::Http { status, .. } if *status == 429 || *status >= 500
841 ) || matches!(err, ProviderError::Transport(_))
842 }
843
844 pub fn is_context_overflow(err: &ProviderError) -> bool {
850 let ProviderError::Http {
851 status: 400,
852 message,
853 ..
854 } = err
855 else {
856 return false;
857 };
858 let m = message.to_lowercase();
859 [
860 "too long",
861 "context length",
862 "context window",
863 "tokens exceed",
864 ]
865 .iter()
866 .any(|needle| m.contains(needle))
867 }
868
869 #[cfg(test)]
870 mod tests {
871 use super::*;
872
873 #[test]
874 fn overflow_detection() {
875 let overflow = ProviderError::Http {
876 status: 400,
877 message:
878 r#"{"error":{"message":"prompt is too long: 210000 tokens > 200000 maximum"}}"#
879 .into(),
880 retry_after: None,
881 };
882 assert!(is_context_overflow(&overflow));
883 let oai = ProviderError::Http {
884 status: 400,
885 message: "This model's maximum context length is 128000 tokens".into(),
886 retry_after: None,
887 };
888 assert!(is_context_overflow(&oai));
889 let plain_400 = ProviderError::Http {
890 status: 400,
891 message: "bad schema".into(),
892 retry_after: None,
893 };
894 assert!(!is_context_overflow(&plain_400));
895 }
896
897 #[test]
901 fn retry_after_is_capped() {
902 let hostile = ProviderError::Http {
903 status: 429,
904 message: String::new(),
905 retry_after: Some(3600),
906 };
907 assert_eq!(
908 classify(&hostile, 1),
909 Decision::Retry {
910 delay: RETRY_AFTER_CAP
911 }
912 );
913 }
914
915 #[test]
917 fn budget_is_deep_enough_for_an_overload() {
918 let overload = ProviderError::Http {
919 status: 529,
920 message: String::new(),
921 retry_after: None,
922 };
923 assert_eq!(MAX_ATTEMPTS, 5);
924 for (attempt, secs) in [(1u32, 1u64), (2, 2), (3, 4), (4, 8)] {
925 assert_eq!(
926 classify(&overload, attempt),
927 Decision::Retry {
928 delay: std::time::Duration::from_secs(secs)
929 },
930 "attempt {attempt}"
931 );
932 }
933 assert_eq!(classify(&overload, MAX_ATTEMPTS), Decision::Fatal);
934 }
935
936 #[test]
939 fn jitter_stays_in_range_and_actually_varies() {
940 let base = std::time::Duration::from_secs(8);
941 let mut seen = std::collections::HashSet::new();
942 for _ in 0..500 {
943 let d = with_jitter(base);
944 assert!(d >= base / 2 && d <= base, "{d:?} outside [base/2, base]");
945 seen.insert(d.as_millis());
946 }
947 assert!(
948 seen.len() > 10,
949 "jitter is not varying: {} values",
950 seen.len()
951 );
952 }
953
954 #[test]
957 fn retry_after_parses_seconds_and_http_date() {
958 const WHEN: u64 = 1_445_412_480;
960 assert_eq!(parse_retry_after("120", WHEN), Some(120));
961 assert_eq!(parse_retry_after(" 120 ", WHEN), Some(120));
962 assert_eq!(
963 parse_retry_after("Wed, 21 Oct 2015 07:30:00 GMT", WHEN),
964 Some(120)
965 );
966 assert_eq!(
968 parse_retry_after("Wed, 21 Oct 2015 07:00:00 GMT", WHEN),
969 Some(0)
970 );
971 assert_eq!(parse_retry_after("not a date", WHEN), None);
972 assert_eq!(
974 parse_retry_after("Mon, 29 Feb 2016 00:00:01 GMT", 1_456_704_000),
975 Some(1)
976 );
977 }
978
979 #[test]
980 fn classify_rules() {
981 let overload = ProviderError::Http {
982 status: 529,
983 message: String::new(),
984 retry_after: Some(7),
985 };
986 assert_eq!(
987 classify(&overload, 1),
988 Decision::Retry {
989 delay: std::time::Duration::from_secs(7)
990 }
991 );
992 assert_eq!(classify(&overload, MAX_ATTEMPTS), Decision::Fatal);
993 let auth = ProviderError::Auth("bad".into());
994 assert_eq!(classify(&auth, 1), Decision::Fatal);
995 assert!(!is_availability(&auth));
996 let transport = ProviderError::Transport("reset".into());
997 assert_eq!(
998 classify(&transport, 2),
999 Decision::Retry {
1000 delay: std::time::Duration::from_secs(2)
1001 }
1002 );
1003 assert!(is_availability(&transport));
1004 let bad_req = ProviderError::Http {
1005 status: 400,
1006 message: String::new(),
1007 retry_after: None,
1008 };
1009 assert_eq!(classify(&bad_req, 1), Decision::Fatal);
1010 }
1011 }
1012}
1013
1014pub mod transform {
1019 use serde_json::Value;
1020
1021 pub fn strip_foreign_reasoning(blocks: &[Value]) -> Vec<Value> {
1024 blocks
1025 .iter()
1026 .filter(|b| {
1027 !matches!(
1028 b.get("type").and_then(Value::as_str),
1029 Some("thinking") | Some("redacted_thinking")
1030 )
1031 })
1032 .cloned()
1033 .collect()
1034 }
1035
1036 #[cfg(test)]
1037 mod tests {
1038 use super::*;
1039 use serde_json::json;
1040
1041 #[test]
1042 fn strips_thinking_keeps_rest() {
1043 let blocks = vec![
1044 json!({"type":"thinking","thinking":"x","signature":"s"}),
1045 json!({"type":"redacted_thinking","data":"d"}),
1046 json!({"type":"text","text":"hi"}),
1047 json!({"type":"tool_use","id":"1","name":"read","input":{}}),
1048 ];
1049 let out = strip_foreign_reasoning(&blocks);
1050 assert_eq!(out.len(), 2);
1051 assert_eq!(out[0]["type"], "text");
1052 assert_eq!(out[1]["type"], "tool_use");
1053 }
1054 }
1055}
1056
1057pub mod repair {
1062 use serde_json::Value;
1063
1064 pub fn parse_or_repair(raw: &str) -> Option<Value> {
1067 if let Ok(v) = serde_json::from_str(raw) {
1068 return Some(v);
1069 }
1070 let without_commas = strip_trailing_commas(raw);
1071 if let Ok(v) = serde_json::from_str(&without_commas) {
1072 return Some(v);
1073 }
1074 serde_json::from_str(&close_truncation(&without_commas)).ok()
1075 }
1076
1077 fn strip_trailing_commas(s: &str) -> String {
1079 let mut out = String::with_capacity(s.len());
1080 let mut in_string = false;
1081 let mut escaped = false;
1082 for c in s.chars() {
1083 if in_string {
1084 out.push(c);
1085 if escaped {
1086 escaped = false;
1087 } else if c == '\\' {
1088 escaped = true;
1089 } else if c == '"' {
1090 in_string = false;
1091 }
1092 continue;
1093 }
1094 match c {
1095 '"' => {
1096 in_string = true;
1097 out.push(c);
1098 }
1099 '}' | ']' => {
1100 while out.ends_with(char::is_whitespace) || out.ends_with(',') {
1101 if out.ends_with(',') {
1102 out.pop();
1103 break;
1104 }
1105 out.pop();
1106 }
1107 out.push(c);
1108 }
1109 _ => out.push(c),
1110 }
1111 }
1112 out
1113 }
1114
1115 fn close_truncation(s: &str) -> String {
1117 let mut stack = Vec::new();
1118 let mut in_string = false;
1119 let mut escaped = false;
1120 for c in s.chars() {
1121 if in_string {
1122 if escaped {
1123 escaped = false;
1124 } else if c == '\\' {
1125 escaped = true;
1126 } else if c == '"' {
1127 in_string = false;
1128 }
1129 continue;
1130 }
1131 match c {
1132 '"' => in_string = true,
1133 '{' => stack.push('}'),
1134 '[' => stack.push(']'),
1135 '}' | ']' => {
1136 stack.pop();
1137 }
1138 _ => {}
1139 }
1140 }
1141 let mut out = s.to_string();
1142 if in_string {
1143 out.push('"');
1144 }
1145 while let Some(closer) = stack.pop() {
1146 out.push(closer);
1147 }
1148 out
1149 }
1150
1151 #[cfg(test)]
1152 mod tests {
1153 use super::*;
1154
1155 #[test]
1156 fn repairs_common_damage_and_rejects_garbage() {
1157 assert_eq!(
1158 parse_or_repair(r#"{"path": "a.rs"}"#).unwrap()["path"],
1159 "a.rs"
1160 );
1161 assert_eq!(
1162 parse_or_repair(r#"{"path": "a.rs",}"#).unwrap()["path"],
1163 "a.rs"
1164 );
1165 assert_eq!(
1166 parse_or_repair(r#"{"items": [1, 2,]}"#).unwrap()["items"][1],
1167 2
1168 );
1169 let v = parse_or_repair(r#"{"command": "cargo tes"#).unwrap();
1171 assert_eq!(v["command"], "cargo tes");
1172 assert_eq!(parse_or_repair(r#"{"t": "a,}"}"#).unwrap()["t"], "a,}");
1174 let v = parse_or_repair(r#"{"t": "say \"hi\"",}"#).unwrap();
1176 assert_eq!(v["t"], "say \"hi\"");
1177 assert!(parse_or_repair("not json at all").is_none());
1178 }
1179 }
1180}
1181
1182pub mod api_error;
1183pub mod catalog;
1184pub mod key;
1185
1186pub const MAX_BLOCK_INDEX: usize = 512;
1190
1191pub trait SseAssembler {
1194 fn handle(&mut self, data: &str) -> Result<Vec<StreamEvent>, ProviderError>;
1195 fn finish(self) -> Result<StreamEvent, ProviderError>;
1196}
1197
1198pub fn drive_sse<B, E, A>(
1206 bytes: B,
1207 mut assembler: A,
1208 idle: std::time::Duration,
1209) -> impl futures_util::Stream<Item = Result<StreamEvent, ProviderError>>
1210where
1211 B: futures_util::Stream<Item = Result<bytes::Bytes, E>>,
1212 E: std::fmt::Display,
1213 A: SseAssembler,
1214{
1215 async_stream::stream! {
1216 let mut parser = SseParser::default();
1217 futures_util::pin_mut!(bytes);
1218 use futures_util::StreamExt;
1219 loop {
1220 let next = match tokio::time::timeout(idle, bytes.next()).await {
1221 Ok(next) => next,
1222 Err(_) => {
1223 yield Err(ProviderError::Transport(format!(
1224 "stream stalled: no data for {}s. The connection is likely dead \
1225 (a proxy dropped it without closing); retry the request.",
1226 idle.as_secs()
1227 )));
1228 return;
1229 }
1230 };
1231 let Some(chunk) = next else { break };
1232 let chunk = match chunk {
1233 Ok(c) => c,
1234 Err(e) => {
1235 yield Err(ProviderError::Transport(format!("stream interrupted: {e}")));
1236 return;
1237 }
1238 };
1239 let payloads = match parser.feed(&chunk) {
1240 Ok(payloads) => payloads,
1241 Err(e) => { yield Err(e); return; }
1242 };
1243 for data in payloads {
1244 match assembler.handle(&data) {
1245 Ok(events) => for ev in events { yield Ok(ev); },
1246 Err(e) => { yield Err(e); return; }
1247 }
1248 }
1249 }
1250 match parser.finish() {
1251 Ok(payloads) => {
1252 for data in payloads {
1253 match assembler.handle(&data) {
1254 Ok(events) => for ev in events { yield Ok(ev); },
1255 Err(e) => { yield Err(e); return; }
1256 }
1257 }
1258 }
1259 Err(e) => { yield Err(e); return; }
1260 }
1261 yield assembler.finish();
1262 }
1263}
1264
1265#[cfg(test)]
1266mod drive_sse_tests {
1267 use super::*;
1268 use futures_util::StreamExt;
1269 use std::time::Duration;
1270
1271 struct NeverEnds;
1272 impl SseAssembler for NeverEnds {
1273 fn handle(&mut self, _: &str) -> Result<Vec<StreamEvent>, ProviderError> {
1274 Ok(vec![])
1275 }
1276 fn finish(self) -> Result<StreamEvent, ProviderError> {
1277 Err(ProviderError::Parse("unreachable".into()))
1278 }
1279 }
1280
1281 #[tokio::test(start_paused = true)]
1285 async fn a_silent_stream_times_out_instead_of_hanging() {
1286 let stalled = futures_util::stream::pending::<Result<bytes::Bytes, std::io::Error>>();
1287 let s = drive_sse(stalled, NeverEnds, Duration::from_secs(5));
1288 futures_util::pin_mut!(s);
1289 let first = s.next().await.expect("an event, not a hang");
1290 match first {
1291 Err(ProviderError::Transport(m)) => {
1292 assert!(m.contains("stalled"), "{m}");
1293 assert!(m.contains("retry"), "errors are prompts: {m}");
1294 }
1295 other => panic!("expected a transport timeout, got {other:?}"),
1296 }
1297 assert!(
1298 s.next().await.is_none(),
1299 "the stream must end after the timeout"
1300 );
1301 }
1302
1303 #[tokio::test(start_paused = true)]
1307 async fn a_slow_but_live_stream_is_not_cut() {
1308 let chunks = futures_util::stream::unfold(0u32, |n| async move {
1309 if n == 6 {
1310 return None;
1311 }
1312 tokio::time::sleep(Duration::from_secs(4)).await;
1313 Some((
1314 Ok::<_, std::io::Error>(bytes::Bytes::from_static(b":keepalive\n\n")),
1315 n + 1,
1316 ))
1317 });
1318 let s = drive_sse(chunks, NeverEnds, Duration::from_secs(5));
1319 futures_util::pin_mut!(s);
1320 let evs: Vec<_> = s.collect().await;
1321 assert_eq!(evs.len(), 1);
1324 assert!(matches!(evs[0], Err(ProviderError::Parse(_))), "{evs:?}");
1325 }
1326}