Skip to main content

vtcode_llm/providers/gemini/
llm_provider.rs

1use super::helpers::InteractionStreamState;
2use super::*;
3use crate::providers::shared::{StreamAssemblyError, extract_data_payload, next_sse_event};
4
5fn normalize_stream_event(event: LLMStreamEvent, interaction_reasoning: bool) -> Vec<NormalizedStreamEvent> {
6    match event {
7        LLMStreamEvent::Reasoning { delta } if interaction_reasoning => {
8            vec![NormalizedStreamEvent::ReasoningDelta { delta, source: ReasoningSource::ProviderSummary }]
9        }
10        LLMStreamEvent::Completed { response } => normalize_completed_event(response),
11        event => event.into_normalized(),
12    }
13}
14
15fn normalize_completed_event(response: Box<LLMResponse>) -> Vec<NormalizedStreamEvent> {
16    let mut events = Vec::new();
17    if let Some(tool_calls) = response.tool_calls.as_ref() {
18        for tool_call in tool_calls {
19            events.push(NormalizedStreamEvent::ToolCallStart {
20                call_id: tool_call.id.clone(),
21                name: tool_call.tool_name().map(ToOwned::to_owned),
22            });
23            if let Some(arguments) = tool_call
24                .raw_input()
25                .filter(|arguments| !arguments.trim().is_empty() && arguments.trim() != "{}")
26            {
27                events.push(NormalizedStreamEvent::ToolCallDelta {
28                    call_id: tool_call.id.clone(),
29                    delta: arguments.to_string(),
30                });
31            }
32        }
33    }
34    events.extend(LLMStreamEvent::Completed { response }.into_normalized());
35    events
36}
37
38/// Shared transport for `generateContent` and `streamGenerateContent`.
39impl GeminiProvider {
40    async fn post_generate_content(
41        &self,
42        url: &str,
43        body: &GenerateContentRequest,
44    ) -> Result<reqwest::Response, LLMError> {
45        self.http_client
46            .post(url)
47            .header("x-goog-api-key", self.api_key.as_ref())
48            .json(body)
49            .send()
50            .await
51            .map_err(|e| format_network_error("Gemini", &e))
52    }
53
54    /// Send a generate request and return only a successful response.
55    ///
56    /// A stale `cachedContent` name (expired or evicted) is retried once
57    /// without the cache and with the full system instruction/tools resent.
58    /// Both `generate` and `stream` share this path so recovery cannot drift
59    /// between them. The dead slot is dropped before the retry so the next
60    /// turn rebuilds it.
61    async fn send_generate_request_with_cache_recovery(
62        &self,
63        url: &str,
64        gemini_request: &GenerateContentRequest,
65        request: &LLMRequest,
66    ) -> Result<reqwest::Response, LLMError> {
67        let response = self.post_generate_content(url, gemini_request).await?;
68        if response.status().is_success() {
69            return Ok(response);
70        }
71
72        let status = response.status();
73        let error_text = crate::providers::common::read_provider_error_body(response).await;
74        if gemini_request.cached_content.is_none()
75            || !explicit_cache::is_stale_cache_error(status.as_u16(), &error_text)
76        {
77            return Err(Self::handle_http_error(status, &error_text));
78        }
79
80        self.explicit_cache.clear();
81        let full_request = self.convert_to_gemini_request(request)?;
82        let retry = self.post_generate_content(url, &full_request).await?;
83        if retry.status().is_success() {
84            return Ok(retry);
85        }
86
87        let retry_status = retry.status();
88        let retry_error_text = crate::providers::common::read_provider_error_body(retry).await;
89        Err(Self::handle_http_error(retry_status, &retry_error_text))
90    }
91}
92
93#[async_trait]
94impl LLMProvider for GeminiProvider {
95    fn name(&self) -> &str {
96        "gemini"
97    }
98
99    fn supports_streaming(&self) -> bool {
100        true
101    }
102
103    fn supports_non_streaming(&self, _model: &str) -> bool {
104        // Pinned so the stream-timeout fallback cannot silently regress.
105        true
106    }
107
108    fn supports_reasoning(&self, model: &str) -> bool {
109        // Codex-inspired robustness: Setting model_supports_reasoning to false
110        // does NOT disable it for known reasoning models.
111        models::google::REASONING_MODELS.contains(&model)
112            || self
113                .model_behavior
114                .as_ref()
115                .and_then(|b| b.model_supports_reasoning)
116                .unwrap_or(false)
117    }
118
119    fn supports_reasoning_effort(&self, model: &str) -> bool {
120        // Same robustness logic for reasoning effort
121        models::google::REASONING_MODELS.contains(&model)
122            || self
123                .model_behavior
124                .as_ref()
125                .and_then(|b| b.model_supports_reasoning_effort)
126                .unwrap_or(false)
127    }
128
129    fn supports_context_caching(&self, model: &str) -> bool {
130        models::google::CACHING_MODELS.contains(&model)
131    }
132
133    fn effective_context_size(&self, model: &str) -> usize {
134        let fallback = if model.contains("gemini-3.1") {
135            1_048_576
136        } else if model.contains("3") || model.contains("1.5-pro") {
137            2_097_152
138        } else {
139            1_048_576
140        };
141        crate::provider::catalog_context_window("gemini", model, fallback)
142    }
143
144    async fn generate(&self, request: LLMRequest) -> Result<LLMResponse, LLMError> {
145        let model = request.model.clone();
146        if self.should_use_interactions(&request) {
147            let interaction_request = self.convert_to_interaction_request(&request)?;
148            let url = format!("{}/interactions", self.base_url);
149            let response = self
150                .http_client
151                .post(&url)
152                .header("x-goog-api-key", self.api_key.as_ref())
153                .json(&interaction_request)
154                .send()
155                .await
156                .map_err(|e| format_network_error("Gemini", &e))?;
157
158            if !response.status().is_success() {
159                let status = response.status();
160                let error_text = crate::providers::common::read_provider_error_body(response).await;
161                return Err(Self::handle_http_error(status, &error_text));
162            }
163
164            let interaction_response: Interaction =
165                response.json().await.map_err(|e| format_parse_error("Gemini", &e))?;
166
167            return Self::convert_from_interaction_response(interaction_response, model);
168        }
169
170        let mut gemini_request = self.convert_to_gemini_request(&request)?;
171        if let Some(cache_name) = self.ensure_explicit_cache(&request, &gemini_request).await? {
172            gemini_request = self.apply_explicit_cache_to_request(gemini_request, &cache_name);
173        }
174
175        let url = format!("{}/models/{}:generateContent", self.base_url, request.model);
176
177        let response = self
178            .send_generate_request_with_cache_recovery(&url, &gemini_request, &request)
179            .await?;
180
181        let gemini_response: GenerateContentResponse =
182            response.json().await.map_err(|e| format_parse_error("Gemini", &e))?;
183
184        Self::convert_from_gemini_response(gemini_response, model)
185    }
186
187    async fn stream(&self, request: LLMRequest) -> Result<LLMStream, LLMError> {
188        if self.should_use_interactions(&request) {
189            let model = request.model.clone();
190            let interaction_request = self.convert_to_interaction_request(&request)?;
191            let url = format!("{}/interactions?alt=sse", self.base_url);
192            let response = self
193                .http_client
194                .post(&url)
195                .header("x-goog-api-key", self.api_key.as_ref())
196                .json(&interaction_request)
197                .send()
198                .await
199                .map_err(|e| format_network_error("Gemini", &e))?;
200
201            if !response.status().is_success() {
202                let status = response.status();
203                let error_text = crate::providers::common::read_provider_error_body(response).await;
204                return Err(Self::handle_http_error(status, &error_text));
205            }
206
207            let stream = {
208                try_stream! {
209                    let mut body_stream = response.bytes_stream();
210                    let mut buf: Vec<u8> = Vec::new();
211                    let mut offset = 0usize;
212                    let mut decoder = crate::providers::shared::Utf8StreamDecoder::new();
213                    let mut state = InteractionStreamState::default();
214
215                    while let Some(chunk_result) = body_stream.next().await {
216                        let chunk = chunk_result
217                            .map_err(|err| format_network_error("Gemini", &err))?;
218
219                        decoder.push_bytes(&chunk, &mut buf);
220
221                        while let Some(event) = next_sse_event(&buf, &mut offset)
222                            .map_err(|e| {
223                                StreamAssemblyError::InvalidPayload(format!("non-utf-8 stream data: {e}"))
224                                    .into_llm_error("Gemini")
225                            })?
226                        {
227
228                            let Some(data_payload) = extract_data_payload(event) else {
229                                continue;
230                            };
231
232                            let trimmed_payload = data_payload.trim();
233                            if trimmed_payload.is_empty() || trimmed_payload == "[DONE]" {
234                                continue;
235                            }
236
237                            let payload: Value = serde_json::from_str(trimmed_payload)
238                                .map_err(|err| {
239                                    StreamAssemblyError::InvalidPayload(err.to_string())
240                                        .into_llm_error("Gemini")
241                                })?;
242
243                            for stream_event in Self::apply_interaction_stream_payload(&mut state, &payload)? {
244                                yield stream_event;
245                            }
246                        }
247
248                        // Drain the consumed prefix so `buf` stays bounded to
249                        // the unprocessed tail rather than growing for the
250                        // entire stream lifetime.
251                        if offset > 0 {
252                            buf.drain(..offset);
253                            offset = 0;
254                        }
255                    }
256
257                    if !state.completed {
258                        let formatted_error = error_display::format_llm_error(
259                            "Gemini",
260                            "Interactions stream ended without an interaction.complete event",
261                        );
262                        Err(LLMError::Provider {
263                            message: formatted_error,
264                            metadata: None,
265                        })?;
266                    }
267
268                    let response =
269                        Self::finalize_interaction_stream_state(state, model)?;
270                    yield LLMStreamEvent::Completed { response: Box::new(response) };
271                }
272            };
273            return Ok(Box::pin(stream));
274        }
275
276        let model = request.model.clone();
277        let mut gemini_request = self.convert_to_gemini_request(&request)?;
278        if let Some(cache_name) = self.ensure_explicit_cache(&request, &gemini_request).await? {
279            gemini_request = self.apply_explicit_cache_to_request(gemini_request, &cache_name);
280        }
281
282        let url = format!("{}/models/{}:streamGenerateContent", self.base_url, request.model);
283
284        let response = self
285            .send_generate_request_with_cache_recovery(&url, &gemini_request, &request)
286            .await?;
287
288        let (event_tx, event_rx) = mpsc::unbounded_channel::<Result<LLMStreamEvent, LLMError>>();
289        let completion_sender = event_tx.clone();
290
291        let streaming_timeout = self.timeouts.streaming_ceiling_seconds;
292
293        let model_clone = model.clone();
294        tokio::spawn(async move {
295            let config = StreamingConfig::with_total_timeout(streaming_timeout);
296            let mut processor = StreamingProcessor::with_config(config);
297            let event_sender = completion_sender.clone();
298            let mut aggregator = crate::providers::shared::StreamAggregator::new(model_clone.clone());
299
300            let mut on_chunk = |chunk: &str| -> Result<(), StreamingError> {
301                if chunk.is_empty() {
302                    return Ok(());
303                }
304
305                if let Some(delta) = Self::apply_stream_delta(&mut aggregator.content, chunk) {
306                    if delta.is_empty() {
307                        return Ok(());
308                    }
309
310                    for event in aggregator.sanitizer.process_chunk(&delta) {
311                        event_sender.send(Ok(event)).map_err(|_e| StreamingError::StreamingError {
312                            message: "Streaming consumer dropped".to_string(),
313                            partial_content: Some(chunk.to_string()),
314                        })?;
315                    }
316                }
317                Ok(())
318            };
319
320            let result = processor.process_stream(response, &mut on_chunk).await;
321            match result {
322                Ok(mut streaming_response) => {
323                    if streaming_response.candidates.is_empty() && !aggregator.content.trim().is_empty() {
324                        streaming_response.candidates.push(StreamingCandidate {
325                            content: Content {
326                                role: "model".to_string(),
327                                parts: vec![Part::Text {
328                                    text: aggregator.content.clone(),
329                                    thought_signature: None,
330                                }],
331                            },
332                            finish_reason: None,
333                            index: Some(0),
334                        });
335                    }
336
337                    match Self::convert_from_streaming_response(streaming_response, model_clone) {
338                        Ok(mut final_response) => {
339                            let aggregator_response = aggregator.finalize();
340                            if final_response.reasoning.is_none() {
341                                final_response.reasoning = aggregator_response.reasoning;
342                            }
343                            if final_response.content.is_none() {
344                                final_response.content = aggregator_response.content;
345                            }
346
347                            let _ = completion_sender
348                                .send(Ok(LLMStreamEvent::Completed { response: Box::new(final_response) }));
349                        }
350                        Err(err) => {
351                            let _ = completion_sender.send(Err(err));
352                        }
353                    }
354                }
355                Err(error) => {
356                    let mapped = Self::map_streaming_error(error);
357                    let _ = completion_sender.send(Err(mapped));
358                }
359            }
360        });
361
362        drop(event_tx);
363
364        let stream = {
365            let mut receiver = event_rx;
366            try_stream! {
367                while let Some(event) = receiver.recv().await {
368                    yield event?;
369                }
370            }
371        };
372
373        Ok(Box::pin(stream))
374    }
375
376    async fn stream_normalized(&self, request: LLMRequest) -> Result<LLMNormalizedStream, LLMError> {
377        let interaction_reasoning = self.should_use_interactions(&request);
378        let mut legacy_stream = self.stream(request).await?;
379        let stream = try_stream! {
380            while let Some(event) = legacy_stream.next().await {
381                for normalized in normalize_stream_event(event?, interaction_reasoning) {
382                    yield normalized;
383                }
384            }
385        };
386
387        Ok(Box::pin(stream))
388    }
389
390    fn supported_models(&self) -> Vec<String> {
391        models::google::SUPPORTED_MODELS.iter().map(|s| s.to_string()).collect()
392    }
393
394    fn validate_request(&self, request: &LLMRequest) -> Result<(), LLMError> {
395        if GeminiProvider::uses_latest_gemini_api(&request.model) {
396            if request.temperature.is_some() || request.top_p.is_some() || request.top_k.is_some() {
397                tracing::warn!(
398                    model = %request.model,
399                    temperature = ?request.temperature,
400                    top_p = ?request.top_p,
401                    top_k = ?request.top_k,
402                    "Sampling parameters (temperature, top_p, top_k) are deprecated for this Gemini model and will be ignored by the API"
403                );
404            }
405        }
406
407        if request.previous_response_id.is_some() && request.response_store == Some(false) {
408            let formatted_error = error_display::format_llm_error(
409                "Gemini",
410                "Interactions with previous_interaction_id cannot set store=false",
411            );
412            return Err(LLMError::InvalidRequest { message: formatted_error, metadata: None });
413        }
414
415        if !models::google::SUPPORTED_MODELS.iter().any(|m| *m == request.model) {
416            let formatted_error =
417                error_display::format_llm_error("Gemini", &format!("Unsupported model: {}", request.model));
418            return Err(LLMError::InvalidRequest { message: formatted_error, metadata: None });
419        }
420
421        if let Some(max_tokens) = request.max_tokens {
422            let model = request.model.as_str();
423            let max_output_tokens = if model.contains("3") { 65536 } else { 8192 };
424
425            if max_tokens > max_output_tokens {
426                let formatted_error = error_display::format_llm_error(
427                    "Gemini",
428                    &format!(
429                        "Requested max_tokens ({max_tokens}) exceeds model limit ({max_output_tokens}) for {model}"
430                    ),
431                );
432                return Err(LLMError::InvalidRequest { message: formatted_error, metadata: None });
433            }
434        }
435
436        Ok(())
437    }
438}
439
440#[cfg(test)]
441mod tests {
442    use super::{LLMStreamEvent, NormalizedStreamEvent, ReasoningSource, normalize_stream_event};
443    use crate::provider::{LLMResponse, ToolCall};
444
445    #[test]
446    fn interaction_reasoning_is_marked_as_public_summary() {
447        let events = normalize_stream_event(LLMStreamEvent::Reasoning { delta: "summary".to_string() }, true);
448
449        assert!(matches!(
450            events.as_slice(),
451            [NormalizedStreamEvent::ReasoningDelta { delta, source }]
452                if delta == "summary" && *source == ReasoningSource::ProviderSummary
453        ));
454    }
455
456    #[test]
457    fn standard_reasoning_remains_unclassified() {
458        let events = normalize_stream_event(LLMStreamEvent::Reasoning { delta: "trace".to_string() }, false);
459
460        assert!(matches!(
461            events.as_slice(),
462            [NormalizedStreamEvent::ReasoningDelta { delta, source }]
463                if delta == "trace" && *source == ReasoningSource::Unknown
464        ));
465    }
466
467    #[test]
468    fn completed_tool_calls_become_structured_events() {
469        let events = normalize_stream_event(
470            LLMStreamEvent::Completed {
471                response: Box::new(LLMResponse {
472                    tool_calls: Some(vec![ToolCall::function(
473                        "call_1".to_string(),
474                        "search_workspace".to_string(),
475                        "{\"query\":\"vtcode\"}".to_string(),
476                    )]),
477                    ..Default::default()
478                }),
479            },
480            false,
481        );
482
483        assert!(matches!(
484            events.as_slice(),
485            [
486                NormalizedStreamEvent::ToolCallStart { call_id, name },
487                NormalizedStreamEvent::ToolCallDelta { call_id: delta_call_id, delta },
488                NormalizedStreamEvent::Done { .. }
489            ] if call_id == "call_1"
490                && delta_call_id == "call_1"
491                && name.as_deref() == Some("search_workspace")
492                && delta == "{\"query\":\"vtcode\"}"
493        ));
494    }
495}