rig_core/providers/gemini/interactions_api/
streaming.rs1use serde::{Deserialize, Serialize};
2
3use super::interactions_api_types::{
4 Content, ContentDelta, FunctionCallContent, Interaction, InteractionSseEvent, InteractionUsage,
5 Step, TextContent, TextDelta, ThoughtContent, ThoughtSignatureDelta, ThoughtSummaryContent,
6 ThoughtSummaryDelta, map_interaction_status,
7};
8use crate::error::ProviderError;
9use crate::operation::{CallFragment, Completion, Finish, IfMalformed, TextPart};
10use crate::providers::gemini::streaming::shared_parts;
11use crate::providers::internal::thoughts::Thoughts;
12use crate::providers::internal::wire;
13use crate::wire::{Decoder, Flow, Out, WireEvent, WireFrame};
14use serde_json::{Map, Value};
15
16const KNOWN_EVENT_TYPES: &[&str] = &[
19 "interaction.created",
20 "interaction.completed",
21 "interaction.status_update",
22 "step.start",
23 "step.delta",
24 "step.stop",
25 "error",
26];
27
28fn classify_interaction_frame(data: &str) -> WireEvent<InteractionSseEvent> {
30 wire::classify_tagged_frame(data, "event_type", |event_type| {
31 KNOWN_EVENT_TYPES.contains(&event_type)
32 })
33}
34
35const INTERACTION_MARKER_KEYS: &[&str] = &["steps", "status", "usage", "object", "id"];
38
39pub enum InteractionsEvent {
41 Sse(InteractionSseEvent),
43 Whole(Interaction),
45}
46
47fn classify_interactions_frame(data: &str) -> WireEvent<InteractionsEvent> {
52 wire::classify_or_untagged(
53 data,
54 "event_type",
55 |data| classify_interaction_frame(data).map(InteractionsEvent::Sse),
56 |data| {
57 wire::classify_marker_keyed_frame::<Interaction>(data, INTERACTION_MARKER_KEYS)
58 .map(InteractionsEvent::Whole)
59 },
60 )
61}
62
63#[derive(Debug, Serialize, Deserialize, Default, Clone)]
65pub struct StreamingCompletionResponse {
66 pub usage: Option<InteractionUsage>,
67 pub interaction: Option<Interaction>,
68 #[serde(skip_serializing_if = "Option::is_none")]
72 pub model_version: Option<String>,
73}
74
75impl From<&StreamingCompletionResponse> for crate::completion::Usage {
76 fn from(value: &StreamingCompletionResponse) -> crate::completion::Usage {
77 value
78 .usage
79 .as_ref()
80 .map(crate::completion::Usage::from)
81 .unwrap_or_default()
82 }
83}
84
85impl From<StreamingCompletionResponse> for crate::completion::Usage {
86 fn from(value: StreamingCompletionResponse) -> crate::completion::Usage {
87 (&value).into()
88 }
89}
90
91#[derive(Default)]
94pub struct InteractionsDecoder<'id> {
95 thoughts: Thoughts<'id>,
97 text: Option<TextPart<'id>>,
99}
100
101enum Chunk {
103 Thought {
104 text: String,
105 signature: Option<String>,
106 },
107 Text(String),
108 Call {
109 name: String,
110 arguments: Option<Value>,
111 id: Option<String>,
112 },
113 Raw(crate::message::AdditionalParams),
116}
117
118impl<'id> InteractionsDecoder<'id> {
119 fn close_text(&mut self, out: &mut Out<'id, Completion>) {
120 if let Some(part) = self.text.take() {
121 out.close_text(part);
122 }
123 }
124
125 fn write(&mut self, chunk: Chunk, out: &mut Out<'id, Completion>) -> Result<(), ProviderError> {
128 match chunk {
129 Chunk::Thought { text, signature } => {
130 if !text.is_empty() {
131 self.close_text(out);
132 }
133 self.thoughts.fragment(out, &text);
134 if let Some(signature) = signature {
135 self.thoughts.signature(out, signature);
136 }
137 }
138 Chunk::Text(text) => {
139 if text.is_empty() {
140 return Ok(());
141 }
142 self.thoughts.boundary();
143 let part = self.text.get_or_insert_with(|| out.text());
144 out.push_text(part, &text);
145 }
146 Chunk::Call {
147 name,
148 arguments,
149 id,
150 } => {
151 self.thoughts.boundary();
152 self.close_text(out);
153 shared_parts::function_call(
154 out,
155 name,
156 arguments.unwrap_or(Value::Object(Map::new())),
157 id,
158 None,
159 )?;
160 }
161 Chunk::Raw(params) => {
162 self.thoughts.boundary();
163 self.close_text(out);
164 let part = out.text();
165 out.text_params(&part, params);
166 out.close_text(part);
167 }
168 }
169 Ok(())
170 }
171}
172
173impl<'id> Decoder<'id, Completion> for InteractionsDecoder<'id> {
176 type Event = InteractionsEvent;
177
178 fn classify(&self, frame: WireFrame) -> WireEvent<InteractionsEvent> {
179 classify_interactions_frame(&frame.as_str())
180 }
181
182 fn decode(
183 &mut self,
184 event: InteractionsEvent,
185 mut out: Out<'id, Completion>,
186 ) -> Result<Flow, ProviderError> {
187 let event = match event {
188 InteractionsEvent::Sse(event) => event,
189 InteractionsEvent::Whole(interaction) => {
192 for content in interaction.output_contents() {
193 if let Some(chunk) = content_chunk(content) {
194 self.write(chunk, &mut out)?;
195 }
196 }
197 InteractionSseEvent::InteractionCompleted {
198 interaction,
199 event_id: None,
200 }
201 }
202 };
203
204 match event {
205 InteractionSseEvent::StepDelta { index, delta, .. } => match delta {
206 ContentDelta::ArgumentsDelta(arguments_delta) => {
207 let index = index as usize;
208 if let Some(fragment) = arguments_delta.arguments
209 && !out.pending_name(index).is_empty()
210 {
211 out.call_fragment(
212 index,
213 CallFragment {
214 arguments: Some(fragment.as_str()),
215 ..CallFragment::default()
216 },
217 )?;
218 } else {
219 tracing::warn!(
220 step_index = index,
221 "arguments_delta with no open function-call step; dropping fragment"
222 );
223 }
224 }
225 ContentDelta::ThoughtSummary(ThoughtSummaryDelta { content }) => {
226 if let ThoughtSummaryContent::Text(text) = content {
227 self.write(
228 Chunk::Thought {
229 text: text.text,
230 signature: None,
231 },
232 &mut out,
233 )?;
234 }
235 }
236 ContentDelta::ThoughtSignature(ThoughtSignatureDelta { signature }) => {
237 self.thoughts.signature(&mut out, signature);
239 }
240 delta => {
241 if let Some(chunk) = delta_content(delta).and_then(content_chunk) {
242 self.write(chunk, &mut out)?;
243 }
244 }
245 },
246 InteractionSseEvent::StepStart { index, step, .. } => {
247 if let Step::FunctionCall(FunctionCallContent {
248 name: Some(name),
249 arguments,
250 id,
251 }) = step
252 {
253 self.thoughts.boundary();
256 self.close_text(&mut out);
257 let index = index as usize;
258 out.call_fragment(
259 index,
260 CallFragment {
261 id: id.as_deref(),
262 name: Some(name.as_str()),
263 ..CallFragment::default()
264 },
265 )?;
266 if let Some(arguments) = arguments.filter(|arguments| {
270 arguments
271 .as_object()
272 .is_none_or(|object| !object.is_empty())
273 }) {
274 out.announce_pending(index, arguments);
275 }
276 } else {
277 for chunk in step_start_chunks(step) {
278 self.write(chunk, &mut out)?;
279 }
280 }
281 }
282 InteractionSseEvent::StepStop { index, .. } => {
283 out.close_pending(index as usize, IfMalformed::Fail)?;
285 }
286 InteractionSseEvent::InteractionCompleted { interaction, .. } => {
287 let span = tracing::Span::current();
288 span.record("gen_ai.response.id", &interaction.id);
289 if let Some(model) = interaction.model.clone() {
290 span.record("gen_ai.response.model", model);
291 }
292 for index in out.pending_calls() {
294 tracing::debug!(
295 index,
296 "closing a function-call step left open at interaction.completed"
297 );
298 out.close_pending(index, IfMalformed::Fail)?;
299 }
300 self.close_text(&mut out);
301 self.thoughts.close(&mut out, None);
302
303 let model_version = interaction.model.clone();
305 let native = StreamingCompletionResponse {
306 usage: interaction.usage,
307 interaction: Some(interaction),
308 model_version,
309 };
310 out.raw(serde_json::to_value(&native)?);
311 let usage = (&native).into();
312 let interaction = native.interaction.as_ref();
313 let finish_reason = interaction
314 .and_then(|interaction| interaction.status.as_ref())
315 .map(map_interaction_status);
316 let response_id = interaction.map(|interaction| interaction.id.clone());
317 return Ok(out.end(Finish {
318 usage,
319 reason: finish_reason,
320 response_id,
321 model: native.model_version,
322 ..Finish::default()
323 }));
324 }
325 event @ InteractionSseEvent::Error { .. } => {
326 let body = serde_json::to_string(&event).unwrap_or_default();
329 return Err(crate::error::ProviderError::from_provider_body(body));
330 }
331 InteractionSseEvent::InteractionCreated { .. }
332 | InteractionSseEvent::InteractionStatusUpdate { .. } => {}
333 }
334 Ok(Flow::More)
335 }
336}
337
338fn delta_content(delta: ContentDelta) -> Option<Content> {
342 match delta {
343 ContentDelta::Text(TextDelta { text, annotations }) => {
344 text.map(|text| Content::Text(TextContent { text, annotations }))
345 }
346 ContentDelta::FunctionCall(call) => Some(Content::FunctionCall(call)),
347 _ => None,
348 }
349}
350
351fn step_start_chunks(step: Step) -> Vec<Chunk> {
352 match step {
353 Step::ModelOutput { content } => content.into_iter().filter_map(content_chunk).collect(),
355 Step::FunctionCall(call) => content_chunk(Content::FunctionCall(call))
356 .into_iter()
357 .collect(),
358 _ => Vec::new(),
359 }
360}
361
362fn content_chunk(content: Content) -> Option<Chunk> {
365 match content {
366 Content::Text(text) if !text.text.is_empty() => Some(Chunk::Text(text.text)),
367 Content::FunctionCall(FunctionCallContent {
368 name,
369 arguments,
370 id,
371 }) => Some(Chunk::Call {
372 name: name?,
373 arguments,
374 id,
375 }),
376 Content::Thought(ThoughtContent {
381 summary, signature, ..
382 }) => {
383 let text: String = summary
384 .unwrap_or_default()
385 .into_iter()
386 .filter_map(|content| match content {
387 ThoughtSummaryContent::Text(text) => Some(text.text),
388 _ => None,
389 })
390 .collect();
391 if text.is_empty() && signature.is_none() {
392 return None;
393 }
394 Some(Chunk::Thought { text, signature })
395 }
396 image @ Content::Image(_) => crate::message::AdditionalParams::from_entries([(
399 crate::providers::gemini::GEMINI_RAW_CONTENT_KEY,
400 serde_json::json!(image),
401 )])
402 .map(Chunk::Raw),
403 _ => None,
404 }
405}
406
407#[cfg(test)]
408mod tests;