repositories / smith
smith
There are many coding harnesses - but this one is fast
owned by admin
smith-ai/src/openai.rs
Raw//! 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<u64, ToolAccumulator>,
finish_reason: Option<String>,
usage: Option<Usage>,
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<StreamEvent>) -> 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<StreamEvent>) -> Result<()> {
let entries: Vec<ToolAccumulator> = 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<StreamEvent>) -> 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<Vec<StreamEvent>> {
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<Vec<StreamEvent>> {
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<StopReason> {
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<String>) -> SmithError {
SmithError::Provider {
fault: ProviderFault::Protocol {
message: message.into(),
},
}
}
fn parse_usage(usage: &serde_json::Value) -> Option<Usage> {
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<u8> {
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<Fixture> {
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<Fixture> {
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::<Vec<_>>(),
);
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<u8> = 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}");
}
}
}