Skip to main content

rig_core/providers/gemini/
streaming.rs

1//! The decoder of GenerateContent replies, a whole `generateContent` body or
2//! a stream of `streamGenerateContent` chunks, read as REST JSON. The REST,
3//! Vertex AI and gRPC wires share it: each part becomes a block whose
4//! provider item is the part as Gemini sent it.
5//!
6//! ```
7//! use rig_core::providers::gemini::streaming::GenerateContentDecoder;
8//!
9//! let decoder = GenerateContentDecoder::default();
10//! # let _ = decoder;
11//! ```
12
13use serde_json::{Map, Value};
14
15use super::completion::blocked_prompt_error;
16use super::completion::{map_google_finish_reason, usage_of};
17use crate::error::ProviderError;
18use crate::json_utils::Lenient;
19use crate::message::{CallId, DocumentSourceKind, Image, MediaType, MimeType, ToolName};
20use crate::operation::{Block, Completion, Finish};
21use crate::providers::internal::wire;
22use crate::wire::{
23    AdapterEvent, AdapterUsage, AdapterVerdict, Decoder, Flow, ObservationSink, Out, WireCitation,
24    WireEvent, WireFrame,
25};
26
27/// The recognizability markers of a `streamGenerateContent` chunk: every
28/// genuine frame carries `candidates`, `usageMetadata` and/or
29/// `promptFeedback` (a blocked prompt's only chunk may carry nothing but
30/// the feedback), and the service's in-band abort carries only `error`. A
31/// frame with any of them is a chunk. A valid ID-only frame is recognized
32/// separately as metadata; other JSON is `Unknown`.
33const RECOGNIZABLE_CHUNK_KEYS: &[&str] =
34    &["candidates", "usageMetadata", "promptFeedback", "error"];
35
36/// A GenerateContent reply document in REST JSON: the whole
37/// `generateContent` body or one `streamGenerateContent` chunk.
38#[derive(Debug, Default, serde::Deserialize)]
39#[serde(transparent)]
40pub struct GenerateContentChunk(pub Map<String, Value>);
41
42/// Decode GenerateContent replies. A text or thought part continues the
43/// block of the part before it while the kind stays the same, and a part
44/// that carries only a thought signature joins that block; any other part is
45/// a block of its own. The provider's end is held until EOF because
46/// hosted-tool rounds can report intermediate finish reasons. The
47/// candidate's grounding and recitation metadata cite the text blocks when
48/// the reply ends.
49#[derive(Debug, Default)]
50pub struct GenerateContentDecoder {
51    /// The latest `finishReason`.
52    finish: Option<String>,
53    /// The latest `usageMetadata`, as Gemini sent it.
54    usage: Option<Value>,
55    model_version: Option<String>,
56    response_id: Option<String>,
57    /// The text or thought run later parts continue, and whether it is
58    /// thought.
59    open: Option<(usize, bool)>,
60    /// Whether the open run holds a signature.
61    signed: bool,
62    /// The index of the block written last.
63    last: Option<usize>,
64    /// A signature sent alone before any block, which the next block takes.
65    signature: Option<String>,
66    /// The answer text written so far, which grounding segments cite.
67    answer: grounding::AnswerText,
68    /// The text block and byte offset in it of each part of the latest
69    /// chunk; `None` for a part that is not answer text.
70    placed: Vec<Option<(usize, usize)>>,
71    /// Where the part being written put its text.
72    placement: Option<(usize, usize)>,
73    /// The latest `groundingMetadata`'s citations, by block.
74    grounding: Vec<(usize, WireCitation)>,
75    /// How many chunks carried parts: a reply over more than one counts a
76    /// grounding segment's bytes in the whole answer, not in a part.
77    chunks: usize,
78    /// The `citationMetadata` sources of every chunk, by block.
79    recitations: Vec<(usize, WireCitation)>,
80}
81
82impl<'id> Decoder<'id, Completion> for GenerateContentDecoder {
83    type Event = GenerateContentChunk;
84
85    fn classify(&self, frame: WireFrame) -> WireEvent<GenerateContentChunk> {
86        // ID-only metadata must not create an unknown-content truncation tail.
87        if GenerateContentDecoder::is_analysis_only(&frame) {
88            return wire::classify_marker_keyed_frame(&frame.as_str(), &["responseId"]);
89        }
90        wire::classify_marker_keyed_frame(&frame.as_str(), RECOGNIZABLE_CHUNK_KEYS)
91    }
92
93    fn decode(
94        &mut self,
95        GenerateContentChunk(data): GenerateContentChunk,
96        mut out: Out<'id, Completion>,
97    ) -> Result<Flow, ProviderError> {
98        let data = Value::Object(data);
99        // Gemini repeats both on every chunk. Record a value when it changes:
100        // every span record counts against the exporter's attribute limit.
101        let span = tracing::Span::current();
102        if let Some(id) = data.str("responseId").filter(|id| !id.is_empty())
103            && self.response_id.as_deref() != Some(id)
104        {
105            span.record("gen_ai.response.id", id);
106            self.response_id = Some(id.to_owned());
107        }
108        if let Some(model) = data.str("modelVersion").filter(|model| !model.is_empty())
109            && self.model_version.as_deref() != Some(model)
110        {
111            span.record("gen_ai.response.model", model);
112            self.model_version = Some(model.to_owned());
113        }
114        if let Some(usage) = data.get("usageMetadata") {
115            self.usage = Some(usage.clone());
116        }
117        if let Some(error) = data.at("/error") {
118            // An in-band failure is a provider error, not a truncation. Its
119            // code is an HTTP status; only error statuses join the retry policy.
120            let status = error
121                .get("code")
122                .and_then(Value::as_u64)
123                .and_then(|code| u16::try_from(code).ok())
124                .and_then(|code| http::StatusCode::from_u16(code).ok())
125                .filter(|status| status.is_client_error() || status.is_server_error());
126            let body = serde_json::json!({ "error": error }).to_string();
127            return Err(match status {
128                Some(status) => ProviderError::from_http_response(status, body),
129                None => ProviderError::from_provider_body(body),
130            });
131        }
132        if let Some(blocked) = data.get("promptFeedback").and_then(blocked_prompt_error) {
133            return Err(blocked);
134        }
135        // The candidates, their content and its parts hold every block and
136        // the finish, so a wrongly typed one fails the reply.
137        let candidate = match data
138            .get("candidates")
139            .map(|candidates| (candidates, candidates.get(0)))
140        {
141            None | Some((Value::Null, _) | (Value::Array(_), None)) => return Ok(Flow::More),
142            Some((Value::Array(_), Some(candidate @ Value::Object(_)))) => candidate,
143            Some((Value::Array(_), Some(_))) => {
144                return Err(malformed("a candidate that is not an object"));
145            }
146            Some(_) => return Err(malformed("candidates that are not a list")),
147        };
148        // Last one wins: an intermediate `finishReason` is superseded by the
149        // reason the turn actually ended on. Proto3 JSON spells an enum
150        // value its schema does not know as its number.
151        match candidate.get("finishReason") {
152            Some(Value::String(name)) => self.finish = Some(name.clone()),
153            Some(Value::Number(number)) => self.finish = Some(format!("FINISH_REASON_{number}")),
154            _ => {}
155        }
156        let parts = match candidate.get("content") {
157            None | Some(Value::Null) => None,
158            Some(content @ Value::Object(_)) => content.get("parts"),
159            Some(_) => return Err(malformed("candidate content that is not an object")),
160        };
161        match parts {
162            None | Some(Value::Null) => {}
163            Some(Value::Array(parts)) => {
164                self.chunks += 1;
165                self.placed.clear();
166                for part in parts {
167                    self.part(part.clone(), &mut out)?;
168                    self.placed.push(self.placement.take());
169                }
170            }
171            Some(_) => return Err(malformed("candidate parts that are not a list")),
172        }
173        // Each chunk restates the grounding whole; recitation sources add up.
174        if let Some(metadata) = candidate.get("groundingMetadata") {
175            self.grounding =
176                grounding::grounding(metadata, &self.placed, &self.answer, self.chunks > 1);
177        }
178        if let Some(metadata) = candidate.get("citationMetadata") {
179            let recitations = grounding::recitations(metadata, &self.answer);
180            self.recitations.extend(recitations);
181        }
182        // A failure is final: nothing after it is read.
183        use crate::completion::FinishReason::{Length, Stop};
184        let reason = self.finish.as_deref().map(map_google_finish_reason);
185        match (reason, candidate.get("finishReason")) {
186            (None | Some(Stop | Length), _) | (_, None) => Ok(Flow::More),
187            _ => self.end(out),
188        }
189    }
190
191    /// Gemini ends a reply at EOF, not at its first finish reason: a
192    /// hosted-tool round can report one before more content arrives.
193    fn eof(&mut self, out: Out<'id, Completion>) -> Result<Flow, ProviderError> {
194        self.end(out)
195    }
196}
197
198impl GenerateContentDecoder {
199    /// End the reply on the finish reason the candidate holds. Without one
200    /// the reply did not end.
201    fn end(&mut self, mut out: Out<'_, Completion>) -> Result<Flow, ProviderError> {
202        let Some(reason) = self.finish.take() else {
203            return Err(ProviderError::Truncated);
204        };
205        self.close(&mut out)?;
206        let citations = std::mem::take(&mut self.grounding)
207            .into_iter()
208            .chain(std::mem::take(&mut self.recitations));
209        for (index, citation) in citations {
210            out.cite(index, citation);
211        }
212        let usage = self.usage.as_ref().map(usage_of).unwrap_or_default();
213        let model = self.model_version.take();
214        let response_id = self.response_id.take();
215        Ok(out.end(Finish {
216            usage,
217            reason: Some(map_google_finish_reason(&reason)),
218            response_id,
219            model,
220            ..Finish::default()
221        }))
222    }
223
224    /// Finish the open text or thought run: Gemini states each part whole.
225    fn close(&mut self, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
226        self.open
227            .take()
228            .map_or(Ok(()), |(index, _)| out.finish(index))
229    }
230
231    /// Write a block at a fresh index, closing the run before it. A text or
232    /// thought block stays open as the run later parts continue; any other
233    /// is whole.
234    fn open(
235        &mut self,
236        block: Block,
237        item: Value,
238        text: &str,
239        out: &mut Out<'_, Completion>,
240    ) -> Result<(), ProviderError> {
241        self.close(out)?;
242        let index = out.fresh_index();
243        let mut item = item;
244        if let (Some(signature), Some(fields)) = (self.signature.take(), item.as_object_mut()) {
245            fields
246                .entry("thoughtSignature")
247                .or_insert(Value::String(signature));
248        }
249        self.last = Some(index);
250        self.signed = item
251            .get("thoughtSignature")
252            .and_then(Value::as_str)
253            .is_some_and(|signature| !signature.is_empty());
254        if matches!(block, Block::Text) {
255            self.answer.open(index, text);
256            self.placement = Some((index, 0));
257        }
258        let run = match &block {
259            Block::Text => Some(false),
260            Block::Reasoning { .. } => Some(true),
261            _ => None,
262        };
263        out.open(index, block, item)?;
264        out.push(index, text)?;
265        self.open = run.map(|thought| (index, thought));
266        self.open.map_or_else(|| out.finish(index), |_| Ok(()))
267    }
268
269    /// Write one part of the candidate's content. A part this decoder does
270    /// not model is an opaque item that replays; one with no data, or that
271    /// is not an object, is kept as one that does not. A signature alone
272    /// joins the block before it, as Gemini attaches it to the part before,
273    /// or the next block when it comes first.
274    fn part(&mut self, part: Value, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
275        let Value::Object(fields) = &part else {
276            return self.open(Block::Opaque { replay: false }, part, "", out);
277        };
278        let thought = part.bool("thought") == Some(true);
279        let signature = part
280            .str("thoughtSignature")
281            .filter(|signature| !signature.is_empty());
282        if let Some(text) = part.str("text") {
283            // An empty part without a signature carries nothing.
284            if text.is_empty() && signature.is_none() {
285                return Ok(());
286            }
287            // A second signed part starts a block of its own, so each
288            // signature stays with the part that carried it.
289            return match self.open {
290                Some((index, kind)) if kind == thought && !(self.signed && signature.is_some()) => {
291                    self.signed |= signature.is_some();
292                    if !thought {
293                        self.placement = self.answer.push(index, text).map(|at| (index, at));
294                    }
295                    out.push(index, text)?;
296                    out.edit(index, |item| merge_part(item, &part))
297                }
298                _ if thought => self.open(
299                    Block::Reasoning { redacted: false },
300                    part.clone(),
301                    text,
302                    out,
303                ),
304                _ => self.open(Block::Text, part.clone(), text, out),
305            };
306        }
307        if let Some(call) = part.obj("functionCall") {
308            let Ok(name) =
309                ToolName::new(call.get("name").and_then(Value::as_str).unwrap_or_default())
310            else {
311                tracing::warn!("Gemini sent a function call without a name; nothing can answer it");
312                return Ok(());
313            };
314            // Proto3 JSON leaves out an empty `args` Struct.
315            let args = call
316                .get("args")
317                .map_or_else(|| "{}".to_owned(), Value::to_string);
318            let id = CallId::from_wire(call.get("id").and_then(Value::as_str).unwrap_or_default());
319            return self.open(Block::Call { id, name }, part.clone(), &args, out);
320        }
321        if !thought
322            && let (Some(mime_type), Some(data)) =
323                (part.at("/inlineData/mimeType"), part.at("/inlineData/data"))
324            && let (Some(mime_type), Some(data)) = (mime_type.as_str(), data.as_str())
325            && let Some(MediaType::Image(media_type)) = MediaType::from_mime_type(mime_type)
326        {
327            let image = Image {
328                data: DocumentSourceKind::Base64(data.to_owned()),
329                media_type: Some(media_type),
330                detail: None,
331                native: None,
332            };
333            return self.open(Block::Image(image), part.clone(), "", out);
334        }
335        let bare = ["thought", "thoughtSignature", "partMetadata"];
336        let data = fields.keys().any(|key| !bare.contains(&key.as_str()));
337        if !data && let Some(signature) = signature {
338            let Some(index) = self.open.map(|(index, _)| index).or(self.last) else {
339                self.signature = Some(signature.to_owned());
340                return Ok(());
341            };
342            return out.edit(index, |item| {
343                if let Some(item) = item.as_object_mut() {
344                    item.insert("thoughtSignature".to_owned(), Value::from(signature));
345                }
346            });
347        }
348        self.open(Block::Opaque { replay: data }, part.clone(), "", out)
349    }
350
351    /// Whether `frame` carries a response id and nothing else, a repeated
352    /// key included.
353    pub(crate) fn is_analysis_only(frame: &WireFrame) -> bool {
354        #[derive(serde::Deserialize)]
355        #[serde(deny_unknown_fields)]
356        struct ResponseIdOnly {
357            #[serde(rename = "responseId")]
358            _id: String,
359        }
360        matches!(
361            wire::classify_marker_keyed_frame::<ResponseIdOnly>(&frame.as_str(), &["responseId"]),
362            WireEvent::Known(_)
363        )
364    }
365
366    /// Project provider verdicts, usage, response identity, and errors
367    /// before normalization. A field of another type is left out, and a
368    /// payload that is not JSON projects nothing.
369    pub(crate) fn project(payload: &[u8], sink: &mut ObservationSink<'_>) {
370        let Ok(reply) = serde_json::from_slice::<Value>(payload) else {
371            return;
372        };
373        if let Some(usage) = reply.get("usageMetadata").filter(|usage| usage.is_object()) {
374            let count = |key: &str| usage.u64(key);
375            let usage = AdapterUsage {
376                input_tokens: count("promptTokenCount"),
377                output_tokens: count("candidatesTokenCount"),
378                total_tokens: count("totalTokenCount"),
379                cached_input_tokens: count("cachedContentTokenCount"),
380                reasoning_tokens: count("thoughtsTokenCount"),
381                tool_input_tokens: count("toolUsePromptTokenCount"),
382            };
383            sink.emit(AdapterEvent::Usage { usage });
384        }
385        let candidate = reply.arr("candidates").first().unwrap_or(&Value::Null);
386        let scrub = |value: Option<&str>| value.map(|value| sink.scrub(value));
387        let block = reply
388            .at("/promptFeedback/blockReason")
389            .and_then(Value::as_str);
390        let verdict = AdapterVerdict {
391            finish_reason: scrub(candidate.str("finishReason")),
392            block_reason: scrub(block),
393            detail: scrub(candidate.str("finishMessage")),
394            model: scrub(reply.str("modelVersion")),
395        };
396        let response_id = scrub(reply.str("responseId"));
397        sink.provider(verdict, response_id);
398        if let Some(error) = reply.get("error").filter(|error| error.is_object()) {
399            let text = |key: &str| error.str(key).map(str::to_owned);
400            let error = crate::observe::ObservedError {
401                code: error.get("code").cloned(),
402                kind: text("status").or_else(|| text("type")),
403                message: text("message"),
404            };
405            error.emit(sink);
406        }
407    }
408}
409
410fn malformed(what: &str) -> ProviderError {
411    ProviderError::Response(format!("Gemini sent {what}"))
412}
413
414/// Merge a streamed part into the part its block holds so far: the text
415/// appends, a non-empty `thoughtSignature` replaces the one held, and any
416/// other field replaces its namesake.
417fn merge_part(item: &mut Value, part: &Value) {
418    let (Value::Object(held), Value::Object(part)) = (&mut *item, part) else {
419        *item = part.clone();
420        return;
421    };
422    for (key, value) in part {
423        if key == "text"
424            && let (Some(Value::String(text)), Value::String(more)) = (held.get_mut("text"), value)
425        {
426            text.push_str(more);
427        } else if key != "thoughtSignature" || value.as_str() != Some("") {
428            held.insert(key.clone(), value.clone());
429        }
430    }
431}
432
433pub mod document;
434mod grounding;
435
436#[cfg(test)]
437mod tests;