Luigit
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}");
        }
    }
}