//! OpenAI-compatible chat-completions stream adapter (`SMH-SPEC-SPEC0001`, //! Providers). //! //! Translates untyped server-sent-event chunks with fragmented tool-call //! argument strings into the normalized vocabulary. All vendor structure, //! the `[DONE]` terminator, and argument accumulation stay inside this //! module; downstream sees only normalized events. use crate::conformance::StreamDecoder; use smith::error::{ProviderFault, Result, SmithError}; use smith::id::{MessageId, ToolCallId}; use smith::stream::{StopReason, StreamEvent, Usage}; use std::collections::BTreeMap; /// Incremental decoder for one OpenAI-compatible stream. #[derive(Debug)] pub struct OpenAiDecoder { buffer: String, message_id: MessageId, tools: BTreeMap, finish_reason: Option, usage: Option, stopped: bool, terminal: bool, } #[derive(Clone, Debug)] struct ToolAccumulator { call_id: ToolCallId, name: String, arguments: String, } const DONE: &str = "[DONE]"; impl OpenAiDecoder { /// A decoder for a fresh assistant message. #[must_use] pub fn new() -> Self { Self { buffer: String::new(), message_id: MessageId::new(), tools: BTreeMap::new(), finish_reason: None, usage: None, stopped: false, terminal: false, } } fn process_event(&mut self, data: &str, events: &mut Vec) -> Result<()> { if self.terminal { return Err(protocol("event after terminal")); } if data.trim() == DONE { return self.finish_done(events); } let chunk: serde_json::Value = serde_json::from_str(data.trim()) .map_err(|e| protocol(format!("chunk is not JSON: {e}")))?; if let Some(fault) = chunk.get("error") { return Err(map_error(fault)); } if let Some(usage) = chunk.get("usage").filter(|value| value.is_object()) { self.usage = parse_usage(usage); } let Some(choice) = chunk .get("choices") .and_then(serde_json::Value::as_array) .and_then(|choices| choices.first()) else { return Ok(()); }; let delta = choice.get("delta").unwrap_or(&serde_json::Value::Null); if let Some(text) = delta.get("content").and_then(serde_json::Value::as_str) && !text.is_empty() { events.push(StreamEvent::text_delta(self.message_id, text)); } if let Some(text) = delta .get("reasoning_content") .and_then(serde_json::Value::as_str) && !text.is_empty() { events.push(StreamEvent::thinking_delta(self.message_id, text)); } if let Some(calls) = delta .get("tool_calls") .and_then(serde_json::Value::as_array) { for call in calls { self.accumulate_tool_call(call)?; } } if let Some(reason) = choice .get("finish_reason") .and_then(serde_json::Value::as_str) { if self.finish_reason.is_some() { return Err(protocol("finish_reason sent twice")); } self.finish_reason = Some(reason.to_string()); self.flush_tools(events)?; } Ok(()) } fn accumulate_tool_call(&mut self, call: &serde_json::Value) -> Result<()> { let index = call .get("index") .and_then(serde_json::Value::as_u64) .ok_or_else(|| protocol("tool call fragment without index"))?; let entry = self.tools.entry(index).or_insert_with(|| ToolAccumulator { call_id: ToolCallId::new(), name: String::new(), arguments: String::new(), }); if let Some(name) = call .get("function") .and_then(|function| function.get("name")) .and_then(serde_json::Value::as_str) { entry.name = name.to_string(); } if let Some(fragment) = call .get("function") .and_then(|function| function.get("arguments")) .and_then(serde_json::Value::as_str) { entry.arguments.push_str(fragment); } Ok(()) } /// Emit one complete `ToolUse` per accumulated call, in index order. fn flush_tools(&mut self, events: &mut Vec) -> Result<()> { let entries: Vec = self.tools.values().cloned().collect(); for entry in entries { if entry.name.is_empty() { return Err(protocol("tool call without a name")); } let input = if entry.arguments.trim().is_empty() { serde_json::json!({}) } else { serde_json::from_str(&entry.arguments) .map_err(|e| protocol(format!("fragmented tool arguments are not JSON: {e}")))? }; events.push(StreamEvent::tool_use( self.message_id, entry.name, input, entry.call_id, )); } self.tools.clear(); Ok(()) } fn finish_done(&mut self, events: &mut Vec) -> Result<()> { let Some(reason) = self.finish_reason.take() else { return Err(protocol("terminal marker without finish_reason")); }; if !self.tools.is_empty() { return Err(protocol("terminal marker with unfinished tool calls")); } let reason = map_finish_reason(&reason)?; self.terminal = true; self.stopped = true; events.push(StreamEvent::stop(reason, self.usage)); Ok(()) } } impl Default for OpenAiDecoder { fn default() -> Self { Self::new() } } impl StreamDecoder for OpenAiDecoder { fn push(&mut self, bytes: &[u8]) -> Result> { self.buffer.push_str(&String::from_utf8_lossy(bytes)); let mut events = Vec::new(); while let Some((data, rest)) = take_sse_data(&self.buffer) { self.buffer = rest; self.process_event(&data, &mut events)?; } Ok(events) } fn finish(&mut self) -> Result> { let mut events = Vec::new(); if !self.buffer.trim().is_empty() { let data = std::mem::take(&mut self.buffer); self.process_event(data.trim_end(), &mut events)?; } if !self.stopped { return Err(SmithError::Provider { fault: ProviderFault::Incomplete, }); } Ok(events) } } /// Take the first complete SSE `data:` payload plus the remaining buffer. fn take_sse_data(buffer: &str) -> Option<(String, String)> { let mut data = String::new(); let mut consumed = 0; for line in buffer.split_inclusive('\n') { let trimmed = line.trim_end_matches(['\r', '\n']); if trimmed.is_empty() { if data.is_empty() { consumed += line.len(); continue; } let rest = buffer[consumed + line.len()..].to_string(); return Some((data, rest)); } if let Some(payload) = trimmed.strip_prefix("data:") { data.push_str(payload.strip_prefix(' ').unwrap_or(payload)); } consumed += line.len(); } None } fn map_finish_reason(reason: &str) -> Result { match reason { "stop" => Ok(StopReason::EndTurn), "length" => Ok(StopReason::Limit), "tool_calls" | "function_call" => Ok(StopReason::ToolUse), "content_filter" => Ok(StopReason::ContentFilter), other => Err(protocol(format!("unmapped finish_reason {other}"))), } } fn map_error(error: &serde_json::Value) -> SmithError { let fault = match error.get("type").and_then(serde_json::Value::as_str) { Some("rate_limit_error" | "insufficient_quota") => ProviderFault::RateLimit { retry_after_ms: None, }, Some("server_error" | "timeout" | "internal_error" | "api_error") => { ProviderFault::Transient { message: message_of(error), } } Some("invalid_request_error") => ProviderFault::Invalid { field: None, message: message_of(error), }, Some("authentication_error" | "invalid_api_key") => ProviderFault::Authentication { message: message_of(error), }, Some("overloaded_error") => ProviderFault::Overloaded { message: message_of(error), }, _ => ProviderFault::Protocol { message: message_of(error), }, }; SmithError::Provider { fault } } fn message_of(error: &serde_json::Value) -> String { error .get("message") .and_then(serde_json::Value::as_str) .unwrap_or("unknown provider error") .to_string() } fn protocol(message: impl Into) -> SmithError { SmithError::Provider { fault: ProviderFault::Protocol { message: message.into(), }, } } fn parse_usage(usage: &serde_json::Value) -> Option { if !usage.is_object() { return None; } Some(Usage { input_tokens: usage .get("prompt_tokens") .and_then(serde_json::Value::as_u64)?, output_tokens: usage .get("completion_tokens") .and_then(serde_json::Value::as_u64)?, input_cache_hit_tokens: usage .get("prompt_tokens_details") .and_then(|details| details.get("cached_tokens")) .and_then(serde_json::Value::as_u64) .unwrap_or(0), input_cache_write_tokens: 0, }) } #[cfg(test)] #[expect( clippy::unwrap_used, reason = "conformance fixtures assert on invariant shapes" )] mod tests { use super::*; use crate::conformance::{Expected, Fixture, conformance_report, sse_boundary_fixtures}; use smith::stream::StopReason; fn sse(lines: &[&str]) -> Vec { let mut out = String::new(); for line in lines { out.push_str("data: "); out.push_str(line); out.push_str("\n\n"); } out.into_bytes() } fn text_fixture() -> Fixture { Fixture { name: "text plus usage", chunks: vec![sse(&[ r#"{"choices":[{"delta":{"role":"assistant","content":""}}]}"#, r#"{"choices":[{"delta":{"content":"Hello"}}]}"#, r#"{"choices":[{"delta":{"content":" world"}}]}"#, r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#, r#"{"choices":[],"usage":{"prompt_tokens":9,"completion_tokens":2,"prompt_tokens_details":{"cached_tokens":4}}}"#, "[DONE]", ])], expected: vec![ Expected::Text("Hello".to_string()), Expected::Text(" world".to_string()), Expected::Stop { reason: StopReason::EndTurn, usage: Some(Usage { input_tokens: 9, output_tokens: 2, input_cache_hit_tokens: 4, input_cache_write_tokens: 0, }), }, ], } } fn tool_fixture() -> Fixture { Fixture { name: "fragmented tool arguments", chunks: vec![sse(&[ r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"write","arguments":""}}]}}]}"#, r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"path\""}}]}}]}"#, r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":":\"f.txt\",\"content\":\"hi\"}"}}]}}]}"#, r#"{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}"#, "[DONE]", ])], expected: vec![ Expected::Tool { name: "write".to_string(), input: serde_json::json!({"path": "f.txt", "content": "hi"}), }, Expected::Stop { reason: StopReason::ToolUse, usage: None, }, ], } } fn thinking_fixture() -> Fixture { Fixture { name: "reasoning deltas", chunks: vec![sse(&[ r#"{"choices":[{"delta":{"reasoning_content":"ponder"}}]}"#, r#"{"choices":[{"delta":{"content":"answer"}}]}"#, r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#, "[DONE]", ])], expected: vec![ Expected::Thinking("ponder".to_string()), Expected::Text("answer".to_string()), Expected::Stop { reason: StopReason::EndTurn, usage: None, }, ], } } fn refusal_and_filter_fixtures() -> Vec { vec![ Fixture { name: "content filter stop", chunks: vec![ sse(&[r#"{"choices":[{"delta":{},"finish_reason":"content_filter"}]}"#]), sse(&["[DONE]"]), ], expected: vec![Expected::Stop { reason: StopReason::ContentFilter, usage: None, }], }, Fixture { name: "rate limit error body", chunks: vec![sse(&[ r#"{"error":{"message":"quota","type":"insufficient_quota"}}"#, ])], expected: vec![Expected::Fault("PROVIDER_RATE_LIMIT".to_string())], }, Fixture { name: "malformed json chunk", chunks: vec![sse(&["{not json}"])], expected: vec![Expected::Fault("PROVIDER_PROTOCOL".to_string())], }, Fixture { name: "truncated before terminal", chunks: vec![sse(&[r#"{"choices":[{"delta":{"content":"partial"}}]}"#])], expected: vec![ Expected::Text("partial".to_string()), Expected::Fault("PROVIDER_INCOMPLETE".to_string()), ], }, Fixture { name: "data after terminal", chunks: vec![ sse(&[r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#]), sse(&["[DONE]"]), sse(&[r#"{"choices":[{"delta":{"content":"late"}}"#]), sse(&["}"]), ], expected: vec![ Expected::Stop { reason: StopReason::EndTurn, usage: None, }, Expected::Fault("PROVIDER_PROTOCOL".to_string()), ], }, Fixture { name: "terminal without finish reason", chunks: vec![sse(&["[DONE]"])], expected: vec![Expected::Fault("PROVIDER_PROTOCOL".to_string())], }, ] } fn boundary_fixtures() -> Vec { sse_boundary_fixtures( r#"data: {"choices":[{"delta":{"content":"Hi"}}]}"#, &sse(&[ r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#, "[DONE]", ]), &[ Expected::Text("Hi".to_string()), Expected::Stop { reason: StopReason::EndTurn, usage: None, }, ], ) } #[test] fn vendor_error_types_map_to_provider_faults() { use crate::conformance::faults::*; let cases = [ ("rate_limit_error", rate_limit()), ("insufficient_quota", rate_limit()), ("server_error", transient()), ("timeout", transient()), ("internal_error", transient()), ("api_error", transient()), ("invalid_request_error", invalid()), ("authentication_error", authentication()), ("invalid_api_key", authentication()), ("overloaded_error", overloaded()), ("not_a_vendor_type", protocol()), ]; for (kind, fault) in cases { let error = serde_json::json!({"type": kind, "message": "m"}); assert_eq!(map_error(&error), SmithError::Provider { fault }, "{kind}"); } } #[test] fn openai_adapter_satisfies_conformance() { let failures = conformance_report( &|| Box::new(OpenAiDecoder::new()), &[ vec![text_fixture(), tool_fixture(), thinking_fixture()], refusal_and_filter_fixtures(), boundary_fixtures(), ] .into_iter() .flatten() .collect::>(), ); assert!(failures.is_empty(), "{failures:?}"); } #[test] fn chunk_splitting_does_not_change_events() { // Same bytes, split at every single byte boundary point sampled. let whole: Vec = text_fixture().chunks.concat(); let mut one_shot = OpenAiDecoder::new(); let mut events = one_shot.push(&whole).unwrap(); events.extend(one_shot.finish().unwrap()); for cut in [0, 1, 7, 40, whole.len() / 2, whole.len() - 3] { let mut split = OpenAiDecoder::new(); let mut collected = split.push(&whole[..cut]).unwrap(); collected.extend(split.push(&whole[cut..]).unwrap()); collected.extend(split.finish().unwrap()); assert_eq!(collected.len(), events.len(), "cut at {cut}"); } } }