1use std::collections::HashMap;
8use std::sync::Arc;
9
10use async_trait::async_trait;
11use regex::Regex;
12use std::sync::LazyLock;
13
14#[cfg(feature = "gemini-llm")]
15use crate::llm::TokenUsage;
16use crate::llm::{
17 BaseLlm, EnvTokenProvider, GcloudTokenProvider, LlmError, LlmRequest, LlmResponse,
18 TokenProvider,
19};
20use crate::utils::variant::{GoogleLlmVariant, get_google_llm_variant};
21
22#[derive(Default)]
24pub struct GeminiLlmParams {
25 pub model: Option<String>,
29 pub api_key: Option<String>,
31 pub vertexai: Option<bool>,
33 pub project: Option<String>,
35 pub location: Option<String>,
37 pub headers: Option<HashMap<String, String>>,
39 pub token_provider: Option<Arc<dyn TokenProvider>>,
41}
42
43pub struct GeminiLlm {
49 model: String,
50 variant: GoogleLlmVariant,
51 #[allow(dead_code)]
53 params: GeminiLlmParams,
54 #[allow(dead_code)]
56 token_provider: Arc<dyn TokenProvider>,
57 #[cfg(feature = "gemini-llm")]
59 client: gemini_genai_rs::Client,
60}
61
62#[cfg(feature = "gemini-llm")]
64fn wire_name(value: &impl serde::Serialize) -> String {
65 serde_json::to_value(value)
66 .ok()
67 .and_then(|v| v.as_str().map(str::to_owned))
68 .unwrap_or_default()
69}
70
71#[cfg(feature = "gemini-llm")]
74fn llm_error(error: gemini_genai_rs::generate::GenerateError) -> LlmError {
75 use gemini_genai_rs::client::http::HttpError;
76 use gemini_genai_rs::generate::GenerateError;
77
78 match error {
79 GenerateError::Http(HttpError::ApiError {
80 status, message, ..
81 }) => LlmError::Api { status, message },
82 GenerateError::Http(HttpError::Request(e)) => LlmError::Transport(e.to_string()),
83 GenerateError::Http(e @ HttpError::RetriesExhausted { .. }) => {
84 LlmError::Transport(e.to_string())
85 }
86 GenerateError::Http(HttpError::Auth(e)) | GenerateError::Auth(e) => {
87 LlmError::Auth(e.to_string())
88 }
89 GenerateError::SafetyBlocked { reason } | GenerateError::PromptBlocked { reason } => {
90 LlmError::ContentFiltered(wire_name(&reason))
91 }
92 other => LlmError::RequestFailed(other.to_string()),
93 }
94}
95
96static SUPPORTED_PATTERNS: LazyLock<Vec<Regex>> = LazyLock::new(|| {
97 vec![
98 Regex::new(r"^gemini-.*$").unwrap(),
99 Regex::new(r"^projects/.*/endpoints/.*$").unwrap(),
100 Regex::new(r"^projects/.*/models/gemini.*$").unwrap(),
101 ]
102});
103
104impl GeminiLlm {
105 pub fn new(mut params: GeminiLlmParams) -> Self {
111 let variant = if let Some(true) = params.vertexai {
113 GoogleLlmVariant::VertexAi
114 } else if let Some(false) = params.vertexai {
115 GoogleLlmVariant::GeminiApi
116 } else {
117 get_google_llm_variant()
118 };
119
120 let model = params
125 .model
126 .clone()
127 .or_else(|| {
128 ["GEMINI_TEXT_MODEL", "GEMINI_MODEL"]
129 .iter()
130 .find_map(|k| std::env::var(k).ok())
131 .filter(|m| !m.trim().is_empty())
132 })
133 .unwrap_or_else(|| match variant {
134 GoogleLlmVariant::GeminiApi => "gemini-flash-latest".to_string(),
135 GoogleLlmVariant::VertexAi => "gemini-2.5-flash".to_string(),
136 });
137 if model.contains("native-audio") || model.contains("-live-") {
138 tracing::warn!(
139 model = %model,
140 "GeminiLlm resolved a Live (bidi) model name for generateContent; set \
141 GEMINI_TEXT_MODEL (or GeminiLlmParams::model) to a text model — a shared \
142 GEMINI_MODEL pointing at the Live model 404s here"
143 );
144 }
145
146 if params.api_key.is_none() && variant == GoogleLlmVariant::GeminiApi {
148 params.api_key = std::env::var("GOOGLE_GENAI_API_KEY")
151 .or_else(|_| std::env::var("GEMINI_API_KEY"))
152 .or_else(|_| std::env::var("GOOGLE_API_KEY"))
153 .ok();
154 }
155
156 if variant == GoogleLlmVariant::VertexAi {
158 if params.project.is_none() {
159 params.project = std::env::var("GOOGLE_CLOUD_PROJECT").ok();
160 }
161 if params.location.is_none() {
162 params.location = std::env::var("GOOGLE_CLOUD_LOCATION").ok();
163 }
164 }
165
166 let token_provider: Arc<dyn TokenProvider> =
171 params.token_provider.take().unwrap_or_else(|| {
172 if variant == GoogleLlmVariant::VertexAi {
173 Arc::new(GcloudTokenProvider::new(std::time::Duration::from_secs(
174 45 * 60,
175 )))
176 } else {
177 Arc::new(EnvTokenProvider)
178 }
179 });
180
181 #[cfg(feature = "gemini-llm")]
186 let client = {
187 use gemini_genai_rs::{Client, prelude::ModelId};
188 match variant {
189 GoogleLlmVariant::GeminiApi => {
190 let api_key = params.api_key.as_deref().unwrap_or("");
191 Client::from_api_key(api_key).model(ModelId::new(model.clone()))
192 }
193 GoogleLlmVariant::VertexAi => {
194 let project = params.project.as_deref().unwrap_or("").to_string();
195 let location = params
196 .location
197 .as_deref()
198 .unwrap_or("us-central1")
199 .to_string();
200 let tp = token_provider.clone();
201 Client::from_vertex_refreshable(project, location, move || tp.token())
202 .model(ModelId::new(model.clone()))
203 }
204 }
205 };
206
207 Self {
208 model,
209 variant,
210 params,
211 token_provider,
212 #[cfg(feature = "gemini-llm")]
213 client,
214 }
215 }
216
217 pub fn from_env() -> Result<Self, LlmError> {
236 Self::try_new(GeminiLlmParams::default())
237 }
238
239 pub fn try_new(params: GeminiLlmParams) -> Result<Self, LlmError> {
242 let llm = Self::new(params);
243 llm.check()?;
244 Ok(llm)
245 }
246
247 fn check(&self) -> Result<(), LlmError> {
249 if super::ModelCapabilities::infer_from_id(&self.model).live_bidi {
250 return Err(LlmError::Config(format!(
251 "`{}` is a Live model, which the text API does not serve. Set \
252 GEMINI_TEXT_MODEL (or `GeminiLlmParams::model`) to a text model such as \
253 `gemini-flash-latest`; keep GEMINI_MODEL for the Live session",
254 self.model
255 )));
256 }
257 match self.variant {
258 GoogleLlmVariant::GeminiApi
259 if self
260 .params
261 .api_key
262 .as_deref()
263 .is_none_or(|k| k.trim().is_empty()) =>
264 {
265 Err(LlmError::Auth(
266 "no Gemini API key. Set GEMINI_API_KEY (create one at \
267 https://aistudio.google.com/apikey), or use Vertex AI with \
268 GOOGLE_GENAI_USE_VERTEXAI=true and GOOGLE_CLOUD_PROJECT"
269 .into(),
270 ))
271 }
272 GoogleLlmVariant::VertexAi
273 if self
274 .params
275 .project
276 .as_deref()
277 .is_none_or(|p| p.trim().is_empty()) =>
278 {
279 Err(LlmError::Config(
280 "Vertex AI needs a project. Set GOOGLE_CLOUD_PROJECT (and optionally \
281 GOOGLE_CLOUD_LOCATION), or unset GOOGLE_GENAI_USE_VERTEXAI to use a \
282 Gemini API key"
283 .into(),
284 ))
285 }
286 _ => Ok(()),
287 }
288 }
289
290 #[cfg(feature = "gemini-llm")]
294 fn from_generate_response(
295 response: gemini_genai_rs::generate::GenerateContentResponse,
296 ) -> Result<LlmResponse, LlmError> {
297 use gemini_genai_rs::prelude::{Content, FinishReason, Role};
298
299 if let Some(reason) = response
300 .prompt_feedback
301 .as_ref()
302 .and_then(|f| f.block_reason)
303 {
304 return Err(LlmError::ContentFiltered(wire_name(&reason)));
305 }
306 let candidate = response.candidates.into_iter().next();
307 let finish = candidate.as_ref().and_then(|c| c.finish_reason);
308 let content = candidate.and_then(|c| c.content).unwrap_or(Content {
309 role: Some(Role::Model),
310 parts: vec![],
311 });
312 if content.parts.is_empty()
313 && let Some(
314 reason @ (FinishReason::Safety
315 | FinishReason::Recitation
316 | FinishReason::Blocklist
317 | FinishReason::ProhibitedContent
318 | FinishReason::Spii),
319 ) = finish
320 {
321 return Err(LlmError::ContentFiltered(wire_name(&reason)));
322 }
323
324 let usage = response.usage_metadata.map(|u| {
326 let prompt = u.prompt_token_count.unwrap_or(0);
327 let completion = u
328 .response_token_count
329 .unwrap_or(0)
330 .saturating_add(u.thoughts_token_count.unwrap_or(0));
331 TokenUsage {
332 prompt_tokens: prompt,
333 completion_tokens: completion,
334 total_tokens: u
335 .total_token_count
336 .unwrap_or(prompt.saturating_add(completion)),
337 }
338 });
339
340 Ok(LlmResponse {
341 content,
342 finish_reason: finish.map(|r| wire_name(&r)),
343 usage,
344 })
345 }
346
347 pub fn is_supported(model: &str) -> bool {
349 SUPPORTED_PATTERNS.iter().any(|re| re.is_match(model))
350 }
351
352 pub fn variant(&self) -> GoogleLlmVariant {
354 self.variant
355 }
356
357 #[cfg(feature = "gemini-llm")]
360 fn to_generate_config(
361 mut request: LlmRequest,
362 ) -> (
363 gemini_genai_rs::generate::GenerateContentConfig,
364 Option<gemini_genai_rs::prelude::ModelId>,
365 ) {
366 use gemini_genai_rs::generate::GenerateContentConfig;
367 use gemini_genai_rs::prelude::{GenerationConfig, ModelId, ThinkingConfig};
368
369 let mut config = if request.contents.is_empty() {
370 GenerateContentConfig::from_text("")
371 } else {
372 GenerateContentConfig::from_contents(std::mem::take(&mut request.contents))
373 };
374 if let Some(sys) = request.system_instruction.take() {
375 config = config.system_instruction(&sys);
376 }
377 config.tools = std::mem::take(&mut request.tools);
378
379 let generation = GenerationConfig {
380 temperature: request.temperature,
381 max_output_tokens: request.max_output_tokens,
382 top_p: request.top_p,
383 top_k: request.top_k,
384 stop_sequences: (!request.stop_sequences.is_empty())
385 .then(|| std::mem::take(&mut request.stop_sequences)),
386 thinking_config: request.thinking_budget.map(|budget| ThinkingConfig {
387 thinking_budget: Some(budget),
388 ..ThinkingConfig::default()
389 }),
390 response_mime_type: request.response_mime_type.take(),
391 response_json_schema: request.response_json_schema.take(),
392 ..GenerationConfig::default()
393 };
394 config.generation_config =
395 (generation != GenerationConfig::default()).then_some(generation);
396
397 (config, request.model.take().map(ModelId::new))
398 }
399}
400
401#[async_trait]
402impl BaseLlm for GeminiLlm {
403 fn model_id(&self) -> &str {
404 &self.model
405 }
406
407 async fn generate(&self, request: LlmRequest) -> Result<LlmResponse, LlmError> {
408 #[cfg(feature = "gemini-llm")]
410 {
411 let (config, model) = Self::to_generate_config(request);
412
413 let response = self
414 .client
415 .generate_content_with(config, model.as_ref())
416 .await
417 .map_err(llm_error)?;
418 Self::from_generate_response(response)
419 }
420
421 #[cfg(not(feature = "gemini-llm"))]
422 {
423 let _ = request;
425 Err(LlmError::RequestFailed(
426 "GeminiLlm requires the 'gemini-llm' feature flag \
427 (depends on gemini-live HTTP client)"
428 .into(),
429 ))
430 }
431 }
432
433 #[cfg(feature = "gemini-llm")]
434 async fn generate_stream(&self, request: LlmRequest) -> Result<super::LlmStream, LlmError> {
435 use futures_util::StreamExt;
436
437 let (config, model) = Self::to_generate_config(request);
438 let chunks = self
439 .client
440 .stream_generate_content_with(config, model.as_ref())
441 .await
442 .map_err(llm_error)?;
443 Ok(chunks
444 .map(|chunk| {
445 chunk
446 .map_err(llm_error)
447 .and_then(Self::from_generate_response)
448 })
449 .boxed())
450 }
451
452 async fn warm_up(&self) -> Result<(), LlmError> {
458 #[cfg(feature = "gemini-llm")]
459 {
460 use gemini_genai_rs::generate::GenerateContentConfig;
461 let config = GenerateContentConfig::from_text(".").max_output_tokens(1);
462 let _ = self.client.generate_content_with(config, None).await;
463 }
464 Ok(())
465 }
466}
467
468#[cfg(test)]
469mod tests {
470 use super::*;
471
472 #[cfg(feature = "gemini-llm")]
475 #[test]
476 fn every_request_field_reaches_the_wire() {
477 let request = LlmRequest {
478 model: Some("gemini-2.5-pro".into()),
479 system_instruction: Some("Be brief.".into()),
480 tools: vec![gemini_genai_rs::prelude::Tool::google_search()],
481 temperature: Some(0.2),
482 max_output_tokens: Some(64),
483 top_p: Some(0.9),
484 top_k: Some(20),
485 stop_sequences: vec!["END".into()],
486 thinking_budget: Some(512),
487 response_mime_type: Some("application/json".into()),
488 response_json_schema: Some(serde_json::json!({ "type": "object" })),
489 ..LlmRequest::from_text("hi")
490 };
491 let (config, model) = GeminiLlm::to_generate_config(request);
492 assert_eq!(
493 model
494 .as_ref()
495 .map(gemini_genai_rs::prelude::ModelId::as_str),
496 Some("gemini-2.5-pro")
497 );
498
499 let body = config.to_request_body();
500 let gc = &body["generationConfig"];
501 assert_eq!(gc["temperature"], 0.2_f32 as f64);
502 assert_eq!(gc["maxOutputTokens"], 64);
503 assert_eq!(gc["topP"], 0.9_f32 as f64);
504 assert_eq!(gc["topK"], 20);
505 assert_eq!(gc["stopSequences"], serde_json::json!(["END"]));
506 assert_eq!(gc["thinkingConfig"]["thinkingBudget"], 512);
507 assert_eq!(gc["responseMimeType"], "application/json");
508 assert_eq!(gc["responseJsonSchema"]["type"], "object");
509 assert!(body["tools"][0].get("googleSearch").is_some(), "{body}");
510 assert!(body["systemInstruction"].is_object(), "{body}");
511 }
512
513 #[cfg(feature = "gemini-llm")]
514 fn wire(json: serde_json::Value) -> Result<LlmResponse, LlmError> {
515 GeminiLlm::from_generate_response(serde_json::from_value(json).unwrap())
516 }
517
518 #[cfg(feature = "gemini-llm")]
520 #[test]
521 fn a_blocked_prompt_is_an_error_with_its_reason() {
522 let err =
523 wire(serde_json::json!({ "promptFeedback": { "blockReason": "SAFETY" } })).unwrap_err();
524 assert!(
525 matches!(&err, LlmError::ContentFiltered(r) if r == "SAFETY"),
526 "{err}"
527 );
528 }
529
530 #[cfg(feature = "gemini-llm")]
531 #[test]
532 fn a_withheld_reply_is_an_error_but_a_truncated_one_is_not() {
533 let withheld = wire(serde_json::json!({
534 "candidates": [{ "finishReason": "PROHIBITED_CONTENT" }]
535 }));
536 assert!(matches!(withheld, Err(LlmError::ContentFiltered(r)) if r == "PROHIBITED_CONTENT"));
537
538 let truncated = wire(serde_json::json!({
539 "candidates": [{
540 "content": { "role": "model", "parts": [{ "text": "Once upon" }] },
541 "finishReason": "MAX_TOKENS"
542 }]
543 }))
544 .unwrap();
545 assert_eq!(truncated.text(), "Once upon");
546 assert_eq!(truncated.finish_reason.as_deref(), Some("MAX_TOKENS"));
547 }
548
549 #[cfg(feature = "gemini-llm")]
552 #[test]
553 fn usage_reads_the_rest_field_names() {
554 let response = wire(serde_json::json!({
555 "candidates": [{ "content": { "role": "model", "parts": [{ "text": "hi" }] } }],
556 "usageMetadata": {
557 "promptTokenCount": 7,
558 "candidatesTokenCount": 3,
559 "thoughtsTokenCount": 5,
560 "totalTokenCount": 15
561 }
562 }))
563 .unwrap();
564 assert_eq!(
565 response.usage,
566 Some(TokenUsage {
567 prompt_tokens: 7,
568 completion_tokens: 8,
569 total_tokens: 15,
570 })
571 );
572 }
573
574 #[cfg(feature = "gemini-llm")]
575 #[test]
576 fn transport_errors_keep_their_status() {
577 use gemini_genai_rs::client::http::HttpError;
578 let err = llm_error(gemini_genai_rs::generate::GenerateError::Http(
579 HttpError::ApiError {
580 status: 429,
581 message: "Resource exhausted".into(),
582 body: None,
583 },
584 ));
585 assert!(err.is_rate_limited(), "{err}");
586 assert!(err.to_string().contains("Resource exhausted"), "{err}");
587 }
588
589 #[test]
590 fn check_names_the_missing_key() {
591 let err = GeminiLlm::try_new(GeminiLlmParams {
592 vertexai: Some(false),
593 api_key: Some(" ".into()),
594 model: Some("gemini-flash-latest".into()),
595 ..Default::default()
596 })
597 .err()
598 .expect("a blank key cannot work");
599 assert!(err.is_auth(), "{err}");
600 assert!(err.to_string().contains("GEMINI_API_KEY"), "{err}");
601 }
602
603 #[test]
604 fn check_rejects_a_live_model_for_text() {
605 let err = GeminiLlm::try_new(GeminiLlmParams {
606 vertexai: Some(false),
607 api_key: Some("key".into()),
608 model: Some("gemini-2.5-flash-native-audio-preview-12-2025".into()),
609 ..Default::default()
610 })
611 .err()
612 .expect("a Live model cannot serve generateContent");
613 assert!(err.to_string().contains("GEMINI_TEXT_MODEL"), "{err}");
614 }
615
616 #[test]
617 fn check_needs_a_vertex_project() {
618 let err = GeminiLlm::try_new(GeminiLlmParams {
619 vertexai: Some(true),
620 project: Some(String::new()),
621 model: Some("gemini-2.5-flash".into()),
622 ..Default::default()
623 })
624 .err()
625 .expect("Vertex AI without a project cannot work");
626 assert!(err.to_string().contains("GOOGLE_CLOUD_PROJECT"), "{err}");
627 }
628
629 #[test]
630 fn check_accepts_a_complete_configuration() {
631 assert!(
632 GeminiLlm::try_new(GeminiLlmParams {
633 vertexai: Some(false),
634 api_key: Some("key".into()),
635 model: Some("gemini-flash-latest".into()),
636 ..Default::default()
637 })
638 .is_ok()
639 );
640 }
641
642 #[cfg(feature = "gemini-llm")]
644 #[test]
645 fn an_unconfigured_request_sends_no_generation_config() {
646 let (config, model) = GeminiLlm::to_generate_config(LlmRequest::from_text("hi"));
647 assert!(model.is_none());
648 assert!(config.to_request_body().get("generationConfig").is_none());
649 }
650
651 #[test]
652 fn default_model_is_the_rolling_flash_alias() {
653 let llm = GeminiLlm::new(GeminiLlmParams {
654 vertexai: Some(false),
655 ..Default::default()
656 });
657 let expected = std::env::var("GEMINI_MODEL")
658 .ok()
659 .filter(|m| !m.trim().is_empty())
660 .unwrap_or_else(|| "gemini-flash-latest".to_string());
661 assert_eq!(llm.model_id(), expected);
662 }
663
664 #[test]
665 fn default_model_on_vertex_is_versioned() {
666 if std::env::var("GEMINI_MODEL").is_ok_and(|m| !m.trim().is_empty()) {
667 return; }
669 let llm = GeminiLlm::new(GeminiLlmParams {
670 vertexai: Some(true),
671 ..Default::default()
672 });
673 assert_eq!(llm.model_id(), "gemini-2.5-flash");
674 }
675
676 #[test]
677 fn explicit_model() {
678 let llm = GeminiLlm::new(GeminiLlmParams {
679 model: Some("gemini-2.0-pro".into()),
680 ..Default::default()
681 });
682 assert_eq!(llm.model_id(), "gemini-2.0-pro");
683 }
684
685 #[test]
686 fn variant_from_params_vertex() {
687 let llm = GeminiLlm::new(GeminiLlmParams {
688 vertexai: Some(true),
689 ..Default::default()
690 });
691 assert_eq!(llm.variant(), GoogleLlmVariant::VertexAi);
692 }
693
694 #[test]
695 fn variant_from_params_gemini_api() {
696 let llm = GeminiLlm::new(GeminiLlmParams {
697 vertexai: Some(false),
698 ..Default::default()
699 });
700 assert_eq!(llm.variant(), GoogleLlmVariant::GeminiApi);
701 }
702
703 #[test]
704 fn is_supported_gemini_models() {
705 assert!(GeminiLlm::is_supported("gemini-2.5-flash"));
706 assert!(GeminiLlm::is_supported("gemini-2.0-pro"));
707 assert!(GeminiLlm::is_supported("gemini-1.5-pro-001"));
708 }
709
710 #[test]
711 fn is_supported_non_gemini_models() {
712 assert!(!GeminiLlm::is_supported("gpt-4"));
713 assert!(!GeminiLlm::is_supported("claude-3-opus"));
714 assert!(!GeminiLlm::is_supported("llama-3"));
715 }
716
717 #[test]
718 fn is_supported_vertex_ai_resource_paths() {
719 assert!(GeminiLlm::is_supported(
720 "projects/my-project/endpoints/12345"
721 ));
722 assert!(GeminiLlm::is_supported(
723 "projects/my-project/models/gemini-2.5-flash"
724 ));
725 }
726
727 #[test]
728 fn model_id_returns_correct_string() {
729 let llm = GeminiLlm::new(GeminiLlmParams {
730 model: Some("gemini-2.5-flash-preview-04-17".into()),
731 ..Default::default()
732 });
733 assert_eq!(llm.model_id(), "gemini-2.5-flash-preview-04-17");
734 }
735
736 #[test]
737 fn base_llm_is_object_safe() {
738 fn _assert_object_safe(_: &dyn BaseLlm) {}
739 }
740
741 #[test]
742 fn gemini_llm_is_send_sync() {
743 fn _assert_send_sync<T: Send + Sync>() {}
744 _assert_send_sync::<GeminiLlm>();
745 }
746}