Luigit
repositories / smith

smith

There are many coding harnesses - but this one is fast

owned by admin

smith-ai/src/gemini.rs

Raw
//! Google Gemini stream adapter (`SMH-SPEC-SPEC0001`, Providers).
//!
//! Translates `streamGenerateContent` server-sent chunks, where function
//! calls arrive complete in one part, into the normalized vocabulary.

use crate::conformance::StreamDecoder;
use smith::error::{ProviderFault, Result, SmithError};
use smith::id::{MessageId, ToolCallId};
use smith::stream::{StopReason, StreamEvent, Usage};

/// Incremental decoder for one Gemini stream.
#[derive(Debug)]
pub struct GeminiDecoder {
    buffer: String,
    message_id: MessageId,
    usage: Usage,
    stop_reason: Option<StopReason>,
    stopped: bool,
}

impl GeminiDecoder {
    /// A decoder for a fresh assistant message.
    #[must_use]
    pub fn new() -> Self {
        Self {
            buffer: String::new(),
            message_id: MessageId::new(),
            usage: Usage::new(),
            stop_reason: None,
            stopped: false,
        }
    }

    fn process_chunk(&mut self, data: &str, events: &mut Vec<StreamEvent>) -> Result<()> {
        if self.stopped || self.stop_reason.is_some() {
            return Err(protocol("chunk after terminal"));
        }
        let chunk: serde_json::Value = serde_json::from_str(data.trim())
            .map_err(|e| protocol(format!("chunk is not JSON: {e}")))?;
        if let Some(error) = chunk.get("error") {
            return Err(map_error(error));
        }
        if let Some(usage) = chunk.get("usageMetadata") {
            if let Some(prompt) = usage
                .get("promptTokenCount")
                .and_then(serde_json::Value::as_u64)
            {
                self.usage.input_tokens = prompt;
            }
            if let Some(total) = usage
                .get("candidatesTokenCount")
                .and_then(serde_json::Value::as_u64)
            {
                self.usage.output_tokens = total;
            }
            if let Some(cached) = usage
                .get("cachedContentTokenCount")
                .and_then(serde_json::Value::as_u64)
            {
                self.usage.input_cache_hit_tokens = cached;
            }
        }
        let Some(candidate) = chunk
            .get("candidates")
            .and_then(serde_json::Value::as_array)
            .and_then(|candidates| candidates.first())
        else {
            return Ok(());
        };
        if let Some(parts) = candidate
            .get("content")
            .and_then(|content| content.get("parts"))
            .and_then(serde_json::Value::as_array)
        {
            for part in parts {
                if let Some(text) = part.get("text").and_then(serde_json::Value::as_str)
                    && !text.is_empty()
                {
                    events.push(StreamEvent::text_delta(self.message_id, text));
                }
                if let Some(call) = part.get("functionCall") {
                    // Gemini delivers function calls complete in one part.
                    let name = call
                        .get("name")
                        .and_then(serde_json::Value::as_str)
                        .ok_or_else(|| protocol("functionCall without name"))?;
                    let input = call
                        .get("args")
                        .cloned()
                        .unwrap_or_else(|| serde_json::json!({}));
                    events.push(StreamEvent::tool_use(
                        self.message_id,
                        name,
                        input,
                        ToolCallId::new(),
                    ));
                }
            }
        }
        if let Some(reason) = candidate
            .get("finishReason")
            .and_then(serde_json::Value::as_str)
        {
            self.stop_reason = Some(map_finish_reason(reason)?);
        }
        Ok(())
    }
}

impl Default for GeminiDecoder {
    fn default() -> Self {
        Self::new()
    }
}

impl StreamDecoder for GeminiDecoder {
    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_chunk(&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_chunk(data.trim_end(), &mut events)?;
        }
        match self.stop_reason.take() {
            // Gemini has no terminal marker; the finishReason chunk is terminal.
            Some(reason) => {
                self.stopped = true;
                events.push(StreamEvent::stop(reason, Some(self.usage)));
                Ok(events)
            }
            None => Err(SmithError::Provider {
                fault: ProviderFault::Incomplete,
            }),
        }
    }
}

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),
        "MAX_TOKENS" => Ok(StopReason::Limit),
        "SAFETY" | "PROHIBITED_CONTENT" | "RECITATION" => Ok(StopReason::ContentFilter),
        other => Err(protocol(format!("unmapped finishReason {other}"))),
    }
}

fn map_error(error: &serde_json::Value) -> SmithError {
    let code = error.get("code").and_then(serde_json::Value::as_i64);
    let message = error
        .get("message")
        .and_then(serde_json::Value::as_str)
        .unwrap_or("unknown provider error")
        .to_string();
    let fault = match code {
        Some(429) => ProviderFault::RateLimit {
            retry_after_ms: None,
        },
        Some(400 | 404) => ProviderFault::Invalid {
            field: None,
            message,
        },
        Some(401 | 403) => ProviderFault::Authentication { message },
        Some(503) => ProviderFault::Overloaded { message },
        Some(500 | 504) => ProviderFault::Transient { message },
        _ => ProviderFault::Protocol { message },
    };
    SmithError::Provider { fault }
}

fn protocol(message: impl Into<String>) -> SmithError {
    SmithError::Provider {
        fault: ProviderFault::Protocol {
            message: message.into(),
        },
    }
}

#[cfg(test)]
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 complete_tool_fixture() -> Fixture {
        Fixture {
            name: "text and complete function call",
            chunks: vec![sse(&[
                r#"{"candidates":[{"content":{"parts":[{"text":"Writing"}]}}]}"#,
                r#"{"candidates":[{"content":{"parts":[{"functionCall":{"name":"write","args":{"path":"f","content":"hi"}}}]}}],"usageMetadata":{"promptTokenCount":5,"candidatesTokenCount":9,"cachedContentTokenCount":2}}"#,
                r#"{"candidates":[{"content":{"parts":[{"text":" now"}]}}]}"#,
                r#"{"candidates":[{"content":{},"finishReason":"STOP"}]}"#,
            ])],
            expected: vec![
                Expected::Text("Writing".to_string()),
                Expected::Tool {
                    name: "write".to_string(),
                    input: serde_json::json!({"path": "f", "content": "hi"}),
                },
                Expected::Text(" now".to_string()),
                Expected::Stop {
                    reason: StopReason::EndTurn,
                    usage: Some(Usage {
                        input_tokens: 5,
                        output_tokens: 9,
                        input_cache_hit_tokens: 2,
                        input_cache_write_tokens: 0,
                    }),
                },
            ],
        }
    }

    fn fault_fixtures() -> Vec<Fixture> {
        vec![
            Fixture {
                name: "safety filter stop",
                chunks: vec![sse(&[
                    r#"{"candidates":[{"content":{},"finishReason":"SAFETY"}]}"#,
                ])],
                expected: vec![Expected::Stop {
                    reason: StopReason::ContentFilter,
                    usage: Some(Usage::new()),
                }],
            },
            Fixture {
                name: "quota error",
                chunks: vec![sse(&[r#"{"error":{"code":429,"message":"quota"}}"#])],
                expected: vec![Expected::Fault("PROVIDER_RATE_LIMIT".to_string())],
            },
            Fixture {
                name: "truncated without finish reason",
                chunks: vec![sse(&[
                    r#"{"candidates":[{"content":{"parts":[{"text":"partial"}]}}]}"#,
                ])],
                expected: vec![
                    Expected::Text("partial".to_string()),
                    Expected::Fault("PROVIDER_INCOMPLETE".to_string()),
                ],
            },
            Fixture {
                name: "chunk after terminal",
                chunks: vec![
                    sse(&[r#"{"candidates":[{"content":{},"finishReason":"STOP"}]}"#]),
                    sse(&[r#"{"candidates":[{"content":{"parts":[{"text":"late"}]}}]}"#]),
                ],
                // The terminal Stop is emitted by finish(); the late chunk
                // faults first, so the stream ends in the protocol fault.
                expected: vec![Expected::Fault("PROVIDER_PROTOCOL".to_string())],
            },
        ]
    }

    fn boundary_fixtures() -> Vec<Fixture> {
        sse_boundary_fixtures(
            r#"data: {"candidates":[{"content":{"parts":[{"text":"Hi"}]}}]}"#,
            &sse(&[r#"{"candidates":[{"content":{},"finishReason":"STOP"}]}"#]),
            &[
                Expected::Text("Hi".to_string()),
                Expected::Stop {
                    reason: StopReason::EndTurn,
                    usage: Some(Usage::new()),
                },
            ],
        )
    }

    #[test]
    fn vendor_status_codes_map_to_provider_faults() {
        use crate::conformance::faults::*;
        let cases = [
            (429, rate_limit()),
            (400, invalid()),
            (404, invalid()),
            (401, authentication()),
            (403, authentication()),
            (503, overloaded()),
            (500, transient()),
            (504, transient()),
            (418, protocol()),
        ];
        for (code, fault) in cases {
            let error = serde_json::json!({"code": code, "message": "m"});
            assert_eq!(map_error(&error), SmithError::Provider { fault }, "{code}");
        }
    }

    #[test]
    fn gemini_adapter_satisfies_conformance() {
        let mut fixtures = vec![complete_tool_fixture()];
        fixtures.extend(boundary_fixtures());
        fixtures.extend(fault_fixtures());
        let failures = conformance_report(&|| Box::new(GeminiDecoder::new()), &fixtures);
        assert!(failures.is_empty(), "{failures:?}");
    }
}