use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use futures_util::stream::BoxStream;
use hotl_types::{Item, StopReason, TokenUsage};
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolDef {
pub name: String,
pub description: String,
pub input_schema: Value,
}
#[derive(Debug, Clone)]
pub struct SamplingRequest {
pub model: String,
pub max_tokens: u32,
pub system: Arc<str>,
pub items: Arc<Vec<Item>>,
pub tools: Arc<[ToolDef]>,
pub thinking: bool,
pub cache_static: bool,
pub turn_context: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "event", rename_all = "snake_case")]
pub enum StreamEvent {
Started,
BlockStart {
index: usize,
kind: String,
},
TextDelta {
index: usize,
text: String,
},
ThinkingDelta {
index: usize,
text: String,
},
ToolInputDelta {
index: usize,
json: String,
},
BlockEnd {
index: usize,
},
Retrying {
attempt: u32,
reason: String,
},
Completed {
stop: StopReason,
usage: TokenUsage,
blocks: Vec<Value>,
},
}
#[derive(Debug, thiserror::Error)]
pub enum ProviderError {
#[error("authentication failed: {0}")]
Auth(String),
#[error("HTTP {status}: {message}")]
Http {
status: u16,
message: String,
retry_after: Option<u64>,
},
#[error("transport error: {0}")]
Transport(String),
#[error("stream parse error: {0}")]
Parse(String),
}
pub trait Provider: Send + Sync {
fn stream(
&self,
req: SamplingRequest,
) -> BoxStream<'static, Result<StreamEvent, ProviderError>>;
}
pub struct ScriptedProvider {
scripts: Mutex<VecDeque<Vec<Result<StreamEvent, ProviderError>>>>,
requests: Mutex<Vec<SamplingRequest>>,
}
impl ScriptedProvider {
pub fn new(scripts: Vec<Vec<Result<StreamEvent, ProviderError>>>) -> Self {
Self {
scripts: Mutex::new(scripts.into()),
requests: Mutex::new(Vec::new()),
}
}
pub fn requests(&self) -> Vec<SamplingRequest> {
self.requests.lock().expect("requests mutex").clone()
}
pub fn last_request(&self) -> Option<SamplingRequest> {
self.requests
.lock()
.expect("requests mutex")
.last()
.cloned()
}
pub fn request_count(&self) -> usize {
self.requests.lock().expect("requests mutex").len()
}
pub fn push_script(&self, script: Vec<Result<StreamEvent, ProviderError>>) {
self.scripts
.lock()
.expect("scripted provider mutex")
.push_back(script);
}
pub fn text_reply(text: &str) -> Vec<Result<StreamEvent, ProviderError>> {
vec![
Ok(StreamEvent::Started),
Ok(StreamEvent::BlockStart {
index: 0,
kind: "text".into(),
}),
Ok(StreamEvent::TextDelta {
index: 0,
text: text.into(),
}),
Ok(StreamEvent::BlockEnd { index: 0 }),
Ok(StreamEvent::Completed {
stop: StopReason::EndTurn,
usage: TokenUsage {
input_tokens: 10,
output_tokens: 5,
..Default::default()
},
blocks: vec![serde_json::json!({"type": "text", "text": text})],
}),
]
}
pub fn tool_call(
id: &str,
name: &str,
input: Value,
) -> Vec<Result<StreamEvent, ProviderError>> {
let block = serde_json::json!({"type": "tool_use", "id": id, "name": name, "input": input});
vec![
Ok(StreamEvent::Started),
Ok(StreamEvent::BlockStart {
index: 0,
kind: "tool_use".into(),
}),
Ok(StreamEvent::BlockEnd { index: 0 }),
Ok(StreamEvent::Completed {
stop: StopReason::ToolUse,
usage: TokenUsage {
input_tokens: 10,
output_tokens: 8,
..Default::default()
},
blocks: vec![block],
}),
]
}
}
impl Provider for ScriptedProvider {
fn stream(
&self,
req: SamplingRequest,
) -> BoxStream<'static, Result<StreamEvent, ProviderError>> {
self.requests.lock().expect("requests mutex").push(req);
let script = self
.scripts
.lock()
.expect("scripted provider mutex")
.pop_front()
.unwrap_or_else(|| {
vec![Err(ProviderError::Transport(
"scripted provider exhausted".into(),
))]
});
Box::pin(futures_util::stream::iter(script))
}
}
#[derive(Default)]
pub struct SseParser {
buf: Vec<u8>,
}
const SSE_MAX_BUFFER: usize = 1024 * 1024;
impl SseParser {
pub fn feed(&mut self, chunk: &[u8]) -> Result<Vec<String>, ProviderError> {
self.buf.extend_from_slice(chunk);
let mut out = Vec::new();
let mut start = 0;
while let Some(pos) = self.buf[start..].iter().position(|&b| b == b'\n') {
let line = String::from_utf8_lossy(&self.buf[start..start + pos]);
let line = line.trim_end_matches('\r');
if let Some(data) = line.strip_prefix("data:") {
let data = data.trim_start();
if !data.is_empty() && data != "[DONE]" {
out.push(data.to_string());
}
}
start += pos + 1;
}
self.buf.drain(..start);
if self.buf.len() > SSE_MAX_BUFFER {
return Err(ProviderError::Parse(format!(
"SSE line exceeded {SSE_MAX_BUFFER} bytes without a newline"
)));
}
Ok(out)
}
}
pub mod retry {
use super::ProviderError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Decision {
Retry { after_secs: u64 },
Fatal,
}
pub const MAX_ATTEMPTS: u32 = 3;
pub fn classify(err: &ProviderError, attempt: u32) -> Decision {
if attempt >= MAX_ATTEMPTS {
return Decision::Fatal;
}
match err {
ProviderError::Http {
status,
retry_after,
..
} if *status == 429 || *status >= 500 => Decision::Retry {
after_secs: retry_after.unwrap_or(1u64 << (attempt - 1)),
},
ProviderError::Transport(_) => Decision::Retry {
after_secs: 1u64 << (attempt - 1),
},
_ => Decision::Fatal,
}
}
pub fn is_availability(err: &ProviderError) -> bool {
matches!(
err,
ProviderError::Http { status, .. } if *status == 429 || *status >= 500
) || matches!(err, ProviderError::Transport(_))
}
pub fn is_context_overflow(err: &ProviderError) -> bool {
let ProviderError::Http {
status: 400,
message,
..
} = err
else {
return false;
};
let m = message.to_lowercase();
[
"too long",
"context length",
"context window",
"tokens exceed",
]
.iter()
.any(|needle| m.contains(needle))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn overflow_detection() {
let overflow = ProviderError::Http {
status: 400,
message:
r#"{"error":{"message":"prompt is too long: 210000 tokens > 200000 maximum"}}"#
.into(),
retry_after: None,
};
assert!(is_context_overflow(&overflow));
let oai = ProviderError::Http {
status: 400,
message: "This model's maximum context length is 128000 tokens".into(),
retry_after: None,
};
assert!(is_context_overflow(&oai));
let plain_400 = ProviderError::Http {
status: 400,
message: "bad schema".into(),
retry_after: None,
};
assert!(!is_context_overflow(&plain_400));
}
#[test]
fn classify_rules() {
let overload = ProviderError::Http {
status: 529,
message: String::new(),
retry_after: Some(7),
};
assert_eq!(classify(&overload, 1), Decision::Retry { after_secs: 7 });
assert_eq!(classify(&overload, MAX_ATTEMPTS), Decision::Fatal);
let auth = ProviderError::Auth("bad".into());
assert_eq!(classify(&auth, 1), Decision::Fatal);
assert!(!is_availability(&auth));
let transport = ProviderError::Transport("reset".into());
assert_eq!(classify(&transport, 2), Decision::Retry { after_secs: 2 });
assert!(is_availability(&transport));
let bad_req = ProviderError::Http {
status: 400,
message: String::new(),
retry_after: None,
};
assert_eq!(classify(&bad_req, 1), Decision::Fatal);
}
}
}
pub mod transform {
use serde_json::Value;
pub fn strip_foreign_reasoning(blocks: &[Value]) -> Vec<Value> {
blocks
.iter()
.filter(|b| {
!matches!(
b.get("type").and_then(Value::as_str),
Some("thinking") | Some("redacted_thinking")
)
})
.cloned()
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn strips_thinking_keeps_rest() {
let blocks = vec![
json!({"type":"thinking","thinking":"x","signature":"s"}),
json!({"type":"redacted_thinking","data":"d"}),
json!({"type":"text","text":"hi"}),
json!({"type":"tool_use","id":"1","name":"read","input":{}}),
];
let out = strip_foreign_reasoning(&blocks);
assert_eq!(out.len(), 2);
assert_eq!(out[0]["type"], "text");
assert_eq!(out[1]["type"], "tool_use");
}
}
}
pub mod repair {
use serde_json::Value;
pub fn parse_or_repair(raw: &str) -> Option<Value> {
if let Ok(v) = serde_json::from_str(raw) {
return Some(v);
}
let without_commas = strip_trailing_commas(raw);
if let Ok(v) = serde_json::from_str(&without_commas) {
return Some(v);
}
serde_json::from_str(&close_truncation(&without_commas)).ok()
}
fn strip_trailing_commas(s: &str) -> String {
let mut out = String::with_capacity(s.len());
let mut in_string = false;
let mut escaped = false;
for c in s.chars() {
if in_string {
out.push(c);
if escaped {
escaped = false;
} else if c == '\\' {
escaped = true;
} else if c == '"' {
in_string = false;
}
continue;
}
match c {
'"' => {
in_string = true;
out.push(c);
}
'}' | ']' => {
while out.ends_with(char::is_whitespace) || out.ends_with(',') {
if out.ends_with(',') {
out.pop();
break;
}
out.pop();
}
out.push(c);
}
_ => out.push(c),
}
}
out
}
fn close_truncation(s: &str) -> String {
let mut stack = Vec::new();
let mut in_string = false;
let mut escaped = false;
for c in s.chars() {
if in_string {
if escaped {
escaped = false;
} else if c == '\\' {
escaped = true;
} else if c == '"' {
in_string = false;
}
continue;
}
match c {
'"' => in_string = true,
'{' => stack.push('}'),
'[' => stack.push(']'),
'}' | ']' => {
stack.pop();
}
_ => {}
}
}
let mut out = s.to_string();
if in_string {
out.push('"');
}
while let Some(closer) = stack.pop() {
out.push(closer);
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn repairs_common_damage_and_rejects_garbage() {
assert_eq!(
parse_or_repair(r#"{"path": "a.rs"}"#).unwrap()["path"],
"a.rs"
);
assert_eq!(
parse_or_repair(r#"{"path": "a.rs",}"#).unwrap()["path"],
"a.rs"
);
assert_eq!(
parse_or_repair(r#"{"items": [1, 2,]}"#).unwrap()["items"][1],
2
);
let v = parse_or_repair(r#"{"command": "cargo tes"#).unwrap();
assert_eq!(v["command"], "cargo tes");
assert_eq!(parse_or_repair(r#"{"t": "a,}"}"#).unwrap()["t"], "a,}");
let v = parse_or_repair(r#"{"t": "say \"hi\"",}"#).unwrap();
assert_eq!(v["t"], "say \"hi\"");
assert!(parse_or_repair("not json at all").is_none());
}
}
}
pub mod key;
pub trait SseAssembler {
fn handle(&mut self, data: &str) -> Result<Vec<StreamEvent>, ProviderError>;
fn finish(self) -> Result<StreamEvent, ProviderError>;
}
pub fn drive_sse<B, E, A>(
bytes: B,
mut assembler: A,
) -> impl futures_util::Stream<Item = Result<StreamEvent, ProviderError>>
where
B: futures_util::Stream<Item = Result<bytes::Bytes, E>>,
E: std::fmt::Display,
A: SseAssembler,
{
async_stream::stream! {
let mut parser = SseParser::default();
futures_util::pin_mut!(bytes);
use futures_util::StreamExt;
while let Some(chunk) = bytes.next().await {
let chunk = match chunk {
Ok(c) => c,
Err(e) => {
yield Err(ProviderError::Transport(format!("stream interrupted: {e}")));
return;
}
};
let payloads = match parser.feed(&chunk) {
Ok(payloads) => payloads,
Err(e) => { yield Err(e); return; }
};
for data in payloads {
match assembler.handle(&data) {
Ok(events) => for ev in events { yield Ok(ev); },
Err(e) => { yield Err(e); return; }
}
}
}
yield assembler.finish();
}
}