vtcode_llm/providers/gemini/
llm_provider.rs1use 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
38impl 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 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 true
106 }
107
108 fn supports_reasoning(&self, model: &str) -> bool {
109 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 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 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}