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