gemini_adk_rs/llm/
gemini.rs

1//! Concrete Gemini LLM implementation using gemini-live `Client`.
2//!
3//! The [`GeminiLlm`] struct is always available for type references and registry
4//! wiring. Actual HTTP generation requires the `gemini-llm` feature flag, which
5//! pulls in `gemini-live/http` and `gemini-live/generate`.
6
7use 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/// Parameters for constructing a [`GeminiLlm`].
23#[derive(Default)]
24pub struct GeminiLlmParams {
25    /// Model name. Defaults to the `GEMINI_MODEL` env var if set, else
26    /// `gemini-flash-latest` on Google AI (a rolling alias the catalog
27    /// keeps serving) or `gemini-2.5-flash` on Vertex AI.
28    pub model: Option<String>,
29    /// API key for Gemini API (non-Vertex).
30    pub api_key: Option<String>,
31    /// Whether to use Vertex AI backend.
32    pub vertexai: Option<bool>,
33    /// Google Cloud project ID (Vertex AI only).
34    pub project: Option<String>,
35    /// Google Cloud region (Vertex AI only, defaults to "us-central1").
36    pub location: Option<String>,
37    /// Custom HTTP headers for requests.
38    pub headers: Option<HashMap<String, String>>,
39    /// Custom token provider for VertexAI. Defaults to reading `GOOGLE_ACCESS_TOKEN` env var.
40    pub token_provider: Option<Arc<dyn TokenProvider>>,
41}
42
43/// Concrete Gemini LLM implementation using gemini-live `Client`.
44///
45/// The gemini-live `Client` is created once at construction time and reused for
46/// all `generate()` calls, matching the JS GenAI SDK pattern where a single
47/// `GoogleGenAI` instance is shared across requests.
48pub struct GeminiLlm {
49    model: String,
50    variant: GoogleLlmVariant,
51    /// Stored for constructing the gemini-live `Client` when `gemini-llm` is enabled.
52    #[allow(dead_code)]
53    params: GeminiLlmParams,
54    /// Token provider for VertexAI token refresh.
55    #[allow(dead_code)]
56    token_provider: Arc<dyn TokenProvider>,
57    /// Cached gemini-live Client, created once at construction time.
58    #[cfg(feature = "gemini-llm")]
59    client: gemini_genai_rs::Client,
60}
61
62/// The name the API uses for an enum value (`"MAX_TOKENS"`, not `MaxTokens`).
63#[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/// Keep what the transport knew: the status, whether credentials failed,
72/// whether the provider was reachable at all.
73#[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    /// Create a new `GeminiLlm` from parameters.
106    ///
107    /// Resolves defaults for model, variant, API key, project, and location
108    /// from parameters first, then falls back to environment variables.
109    /// The gemini-live `Client` is created once here and reused for all calls.
110    pub fn new(mut params: GeminiLlmParams) -> Self {
111        // Resolve variant from params or env
112        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        // Resolve model: params, then GEMINI_TEXT_MODEL, then the shared
121        // GEMINI_MODEL, then a per-variant default. Google AI retires dated
122        // names but serves the rolling `gemini-flash-latest` alias; Vertex AI
123        // keeps versioned GA names and does not carry the alias.
124        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        // Resolve API key from params or env
147        if params.api_key.is_none() && variant == GoogleLlmVariant::GeminiApi {
148            // Same acceptance chain as the Live connect path, so one exported
149            // variable works for both halves of the stack.
150            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        // Resolve project/location from env for Vertex AI
157        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        // Resolve token provider for VertexAI.
167        // Default to GcloudTokenProvider (env var -> gcloud CLI fallback) for VertexAI,
168        // matching the auth resolution in build_session_config(). For GeminiApi, use
169        // EnvTokenProvider since API key auth doesn't need token refresh.
170        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        // Create the gemini-live Client once, reuse across generate() calls.
182        // For VertexAI, use from_vertex_refreshable() so the token is dynamically
183        // refreshed on every REST API call (via auth_headers()), preventing 401
184        // errors from stale tokens during long-running sessions.
185        #[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    /// A `GeminiLlm` configured from the environment, checked before any
218    /// request is made.
219    ///
220    /// Reads the same variables as [`new`](Self::new) —
221    /// `GEMINI_API_KEY` (or `GOOGLE_GENAI_API_KEY`, `GOOGLE_API_KEY`) for Google
222    /// AI; `GOOGLE_GENAI_USE_VERTEXAI=true` with `GOOGLE_CLOUD_PROJECT` and
223    /// optionally `GOOGLE_CLOUD_LOCATION` for Vertex AI; `GEMINI_TEXT_MODEL` or
224    /// `GEMINI_MODEL` for the model — and fails with a message naming the
225    /// variable to set when one is missing, or when the model is a Live model
226    /// the text API does not serve. [`new`](Self::new) accepts the same
227    /// configuration and fails on the first request instead.
228    ///
229    /// ```no_run
230    /// use gemini_adk_rs::llm::GeminiLlm;
231    ///
232    /// let llm = GeminiLlm::from_env()?;
233    /// # Ok::<(), gemini_adk_rs::llm::LlmError>(())
234    /// ```
235    pub fn from_env() -> Result<Self, LlmError> {
236        Self::try_new(GeminiLlmParams::default())
237    }
238
239    /// Like [`new`](Self::new), but checks the resolved configuration first;
240    /// see [`from_env`](Self::from_env).
241    pub fn try_new(params: GeminiLlmParams) -> Result<Self, LlmError> {
242        let llm = Self::new(params);
243        llm.check()?;
244        Ok(llm)
245    }
246
247    /// What would make every request fail, found before sending one.
248    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    /// Turn a wire response into an [`LlmResponse`], or the error it carries:
291    /// a blocked prompt or a reply withheld by content safety is an error, not
292    /// an empty answer.
293    #[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        // Thinking tokens are billed as output, so they count as completion.
325        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    /// Check if a model name is supported by `GeminiLlm`.
348    pub fn is_supported(model: &str) -> bool {
349        SUPPORTED_PATTERNS.iter().any(|re| re.is_match(model))
350    }
351
352    /// Get the variant (VertexAI vs GeminiApi).
353    pub fn variant(&self) -> GoogleLlmVariant {
354        self.variant
355    }
356
357    /// Map every field of an [`LlmRequest`] onto the wire request, and the
358    /// per-request model override onto the model to call.
359    #[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        // Feature-gate the actual HTTP call behind gemini-live's generate + http features.
409        #[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            // Suppress unused-variable warnings when the feature is disabled.
424            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    /// Pre-warm the HTTP connection pool by making a lightweight request.
453    ///
454    /// Establishes the TCP+TLS connection so the first real `generate()`
455    /// call doesn't pay the ~100-300ms handshake penalty. reqwest's
456    /// connection pool keeps it alive for subsequent calls.
457    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    /// Every setting an agent can make must reach the wire body; a field that
473    /// is accepted and dropped is the bug this guards against.
474    #[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    /// A blocked prompt used to come back as an empty answer.
519    #[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    /// `generateContent` names the output count `candidatesTokenCount`; it was
550    /// read as zero. Thinking tokens are billed as output.
551    #[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    /// A request with no settings sends no `generationConfig` at all.
643    #[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; // env override wins by design; nothing to assert here
668        }
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}