Skip to main content

agent_sdk_providers/impls/
gemini.rs

1//! Google Gemini API provider implementation.
2//!
3//! This module provides an implementation of `LlmProvider` for the Google Gemini
4//! API (`generativelanguage.googleapis.com`).
5
6pub(crate) mod data;
7
8use crate::attachments::validate_request_attachments;
9use crate::provider::LlmProvider;
10use crate::streaming::{StreamBox, StreamDelta, StreamErrorKind, reqwest_error_delta};
11use agent_sdk_foundation::llm::{ChatOutcome, ChatRequest, ChatResponse, ThinkingConfig};
12use anyhow::Result;
13use async_trait::async_trait;
14use data::{
15    ApiContent, ApiFunctionCallingConfig, ApiGenerateContentRequest, ApiGenerateContentResponse,
16    ApiGenerationConfig, ApiPart, ApiUsageMetadata, build_api_contents, build_content_blocks,
17    convert_tools_to_config, gemini_response_schema, map_finish_reason, map_thinking_config,
18};
19use reqwest::StatusCode;
20
21const API_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta";
22
23/// Connect timeout for the HTTP client (matches Anthropic/Vertex).
24const CONNECT_TIMEOUT_SECS: u64 = 30;
25/// TCP keepalive interval to keep long streaming connections from dropping.
26const TCP_KEEPALIVE_SECS: u64 = 30;
27/// Per-request read timeout for the **non-streaming** `chat()` path. Bounds a
28/// black-holed endpoint so a single turn cannot hang the agent loop forever.
29/// Streaming requests intentionally have no overall timeout.
30const CHAT_READ_TIMEOUT_SECS: u64 = 300;
31
32/// Max page size the Gemini `ListModels` endpoint accepts (default is 50).
33const MODELS_PAGE_SIZE: u32 = 1000;
34/// Upper bound on pages followed by `list_models`, guarding against a server
35/// that never clears `nextPageToken`.
36const MODELS_MAX_PAGES: usize = 100;
37
38/// Build the shared HTTP client with connect + keepalive timeouts, falling back
39/// to a default client (with a logged warning) if the builder fails.
40fn build_http_client() -> reqwest::Client {
41    reqwest::Client::builder()
42        .connect_timeout(std::time::Duration::from_secs(CONNECT_TIMEOUT_SECS))
43        .tcp_keepalive(std::time::Duration::from_secs(TCP_KEEPALIVE_SECS))
44        .build()
45        .unwrap_or_else(|error| {
46            log::warn!(
47                "failed to build Gemini HTTP client with timeouts ({error}); using default client"
48            );
49            reqwest::Client::new()
50        })
51}
52
53// Gemini 3.1 series
54pub const MODEL_GEMINI_31_PRO: &str = "gemini-3.1-pro-preview";
55pub const MODEL_GEMINI_31_FLASH_LITE: &str = "gemini-3.1-flash-lite-preview";
56
57// Gemini 3 series
58pub const MODEL_GEMINI_3_FLASH: &str = "gemini-3-flash-preview";
59
60// Legacy Gemini 3.0 Pro model kept for explicit opt-in.
61pub const MODEL_GEMINI_3_PRO: &str = "gemini-3.0-pro";
62
63// Gemini 2.5 series
64pub const MODEL_GEMINI_25_FLASH: &str = "gemini-2.5-flash";
65pub const MODEL_GEMINI_25_PRO: &str = "gemini-2.5-pro";
66
67// Gemini 2.0 series
68pub const MODEL_GEMINI_2_FLASH: &str = "gemini-2.0-flash";
69pub const MODEL_GEMINI_2_FLASH_LITE: &str = "gemini-2.0-flash-lite";
70
71/// Google Gemini LLM provider.
72#[derive(Clone)]
73pub struct GeminiProvider {
74    client: reqwest::Client,
75    api_key: String,
76    model: String,
77    base_url: String,
78    thinking: Option<ThinkingConfig>,
79    /// When true, send the API key via `x-goog-api-key` header instead of a
80    /// query parameter. Required when routing through proxies.
81    use_header_auth: bool,
82    /// Extra headers applied to every request (e.g. for gateway authentication).
83    extra_headers: Vec<(String, String)>,
84}
85
86impl GeminiProvider {
87    /// The conventional environment variable holding the Gemini API key.
88    pub const API_KEY_ENV: &'static str = "GEMINI_API_KEY";
89
90    /// Create a new Gemini provider with the specified API key and model.
91    #[must_use]
92    pub fn new(api_key: impl Into<String>, model: impl Into<String>) -> Self {
93        Self {
94            client: build_http_client(),
95            api_key: api_key.into(),
96            model: model.into(),
97            base_url: API_BASE_URL.to_owned(),
98            thinking: None,
99            use_header_auth: true,
100            extra_headers: Vec::new(),
101        }
102    }
103
104    /// Effective output-token budget for a request.
105    ///
106    /// Mirrors the Anthropic provider: when the caller did not explicitly set
107    /// `max_tokens`, substitute the provider/model default
108    /// ([`default_max_tokens`](LlmProvider::default_max_tokens)) instead of
109    /// silently capping at `ChatRequest::DEFAULT_MAX_TOKENS`.
110    fn effective_max_tokens(&self, request: &ChatRequest) -> u32 {
111        if request.max_tokens_explicit {
112            request.max_tokens
113        } else {
114            self.default_max_tokens()
115        }
116    }
117
118    /// Create a provider using Gemini Flash, reading the API key from the
119    /// conventional [`GEMINI_API_KEY`](Self::API_KEY_ENV) environment variable.
120    ///
121    /// # Panics
122    ///
123    /// Panics if `GEMINI_API_KEY` is not set. Prefer
124    /// [`try_from_env`](Self::try_from_env) outside of examples/tests.
125    #[must_use]
126    pub fn from_env() -> Self {
127        Self::try_from_env().unwrap_or_else(|e| panic!("{e}"))
128    }
129
130    /// Create a provider using Gemini Flash, reading the API key from the
131    /// conventional [`GEMINI_API_KEY`](Self::API_KEY_ENV) environment variable.
132    ///
133    /// # Errors
134    ///
135    /// Returns an error if `GEMINI_API_KEY` is unset or not valid UTF-8.
136    pub fn try_from_env() -> Result<Self> {
137        let api_key = std::env::var(Self::API_KEY_ENV).map_err(|_| {
138            anyhow::anyhow!("environment variable `{}` is not set", Self::API_KEY_ENV)
139        })?;
140        Ok(Self::flash(api_key))
141    }
142
143    /// Create a provider using Gemini 3 Flash Preview (fast and capable, current default).
144    #[must_use]
145    pub fn flash(api_key: impl Into<String>) -> Self {
146        Self::new(api_key, MODEL_GEMINI_3_FLASH)
147    }
148
149    /// Create a provider using Gemini 3.1 Flash Lite Preview.
150    #[must_use]
151    pub fn flash_lite_31(api_key: String) -> Self {
152        Self::new(api_key, MODEL_GEMINI_31_FLASH_LITE.to_owned())
153    }
154
155    /// Create a provider using Gemini 2.0 Flash Lite (fastest, most cost-effective).
156    #[must_use]
157    pub fn flash_lite(api_key: String) -> Self {
158        Self::new(api_key, MODEL_GEMINI_2_FLASH_LITE.to_owned())
159    }
160
161    /// Create a provider using Gemini 3.1 Pro Preview.
162    #[must_use]
163    pub fn pro_31(api_key: String) -> Self {
164        Self::new(api_key, MODEL_GEMINI_31_PRO.to_owned())
165    }
166
167    /// Create a provider using Gemini 3.1 Pro Preview (current recommended pro model).
168    #[must_use]
169    pub fn pro(api_key: String) -> Self {
170        Self::new(api_key, MODEL_GEMINI_31_PRO.to_owned())
171    }
172
173    /// Set the provider-owned thinking configuration for this model.
174    #[must_use]
175    pub const fn with_thinking(mut self, thinking: ThinkingConfig) -> Self {
176        self.thinking = Some(thinking);
177        self
178    }
179
180    /// Override the base URL.
181    #[must_use]
182    pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
183        self.base_url = base_url.into();
184        self
185    }
186
187    /// Send the API key via `x-goog-api-key` header instead of `?key=` query
188    /// parameter. Required when routing through proxies.
189    #[must_use]
190    pub const fn with_header_auth(mut self) -> Self {
191        self.use_header_auth = true;
192        self
193    }
194
195    /// Add extra HTTP headers applied to every request.
196    #[must_use]
197    pub fn with_extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
198        self.extra_headers = headers;
199        self
200    }
201
202    /// Apply auth + extra headers. Skips provider auth when `api_key` is
203    /// empty (BYOK gateway mode).
204    fn apply_auth(&self, builder: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
205        let builder = if self.api_key.is_empty() {
206            builder
207        } else if self.use_header_auth {
208            builder.header("x-goog-api-key", &self.api_key)
209        } else {
210            builder.query(&[("key", &self.api_key)])
211        };
212        self.extra_headers
213            .iter()
214            .fold(builder, |b, (k, v)| b.header(k.as_str(), v.as_str()))
215    }
216}
217
218#[async_trait]
219#[allow(clippy::too_many_lines)]
220impl LlmProvider for GeminiProvider {
221    async fn chat(&self, request: ChatRequest) -> Result<ChatOutcome> {
222        let thinking = match self.resolve_thinking_config(request.thinking.as_ref()) {
223            Ok(thinking) => thinking,
224            Err(error) => return Ok(ChatOutcome::InvalidRequest(error.to_string())),
225        };
226        if let Err(error) = validate_request_attachments(self.provider(), self.model(), &request) {
227            return Ok(ChatOutcome::InvalidRequest(error.to_string()));
228        }
229        let contents = build_api_contents(&request.messages);
230        let tools = request
231            .tools
232            .as_ref()
233            .map(|t| convert_tools_to_config(t.clone()));
234        let tool_config = request
235            .tool_choice
236            .as_ref()
237            .map(ApiFunctionCallingConfig::from_tool_choice);
238        let system_instruction = if request.system.is_empty() {
239            None
240        } else {
241            Some(ApiContent {
242                role: None,
243                parts: vec![ApiPart::Text {
244                    text: request.system.clone(),
245                    thought_signature: None,
246                }],
247            })
248        };
249
250        let thinking_config = thinking.as_ref().map(map_thinking_config);
251        let (response_mime_type, response_schema) =
252            request.response_format.as_ref().map_or((None, None), |rf| {
253                (
254                    Some("application/json"),
255                    Some(gemini_response_schema(&rf.schema)),
256                )
257            });
258
259        let max_tokens = self.effective_max_tokens(&request);
260        let api_request = ApiGenerateContentRequest {
261            contents: &contents,
262            system_instruction: system_instruction.as_ref(),
263            tools: tools.as_ref().map(std::slice::from_ref),
264            tool_config,
265            generation_config: Some(ApiGenerationConfig {
266                max_output_tokens: Some(max_tokens),
267                thinking_config,
268                response_mime_type,
269                response_schema,
270            }),
271            cached_content: request.cached_content.as_deref(),
272        };
273
274        log::debug!(
275            "Gemini LLM request model={} max_tokens={}",
276            self.model,
277            max_tokens
278        );
279
280        let builder = self
281            .client
282            .post(format!(
283                "{}/models/{}:generateContent",
284                self.base_url, self.model
285            ))
286            .header("Content-Type", "application/json")
287            .timeout(std::time::Duration::from_secs(CHAT_READ_TIMEOUT_SECS));
288        let response = self
289            .apply_auth(builder)
290            .json(&api_request)
291            .send()
292            .await
293            .map_err(|e| anyhow::anyhow!("request failed: {e}"))?;
294
295        let status = response.status();
296        // Read `Retry-After` off the 429 response before the body is consumed.
297        let retry_after = if status == StatusCode::TOO_MANY_REQUESTS {
298            crate::http::retry_after_from_headers(response.headers())
299        } else {
300            None
301        };
302        let bytes = response
303            .bytes()
304            .await
305            .map_err(|e| anyhow::anyhow!("failed to read response body: {e}"))?;
306
307        log::debug!(
308            "Gemini LLM response status={} body_len={}",
309            status,
310            bytes.len()
311        );
312
313        if status == StatusCode::TOO_MANY_REQUESTS {
314            let retry_after = retry_after.or_else(|| {
315                crate::retry_hints::google_retry_delay(&String::from_utf8_lossy(&bytes))
316            });
317            return Ok(ChatOutcome::RateLimited(retry_after));
318        }
319
320        if status.is_server_error() {
321            let body = String::from_utf8_lossy(&bytes);
322            log::error!("Gemini server error status={status} body={body}");
323            return Ok(ChatOutcome::ServerError(body.into_owned()));
324        }
325
326        if status.is_client_error() {
327            let body = String::from_utf8_lossy(&bytes);
328            log::warn!("Gemini client error status={status} body={body}");
329            return Ok(ChatOutcome::InvalidRequest(body.into_owned()));
330        }
331
332        let api_response: ApiGenerateContentResponse = serde_json::from_slice(&bytes)
333            .map_err(|e| anyhow::anyhow!("failed to parse response: {e}"))?;
334
335        let candidate = api_response
336            .candidates
337            .into_iter()
338            .next()
339            .ok_or_else(|| anyhow::anyhow!("no candidates in response"))?;
340
341        let content = build_content_blocks(&candidate.content);
342
343        if content.is_empty() && !candidate.content.parts.is_empty() {
344            log::warn!(
345                "Gemini parts not converted to content blocks raw_parts={:?}",
346                candidate.content.parts
347            );
348        }
349
350        let has_tool_calls = content
351            .iter()
352            .any(|b| matches!(b, agent_sdk_foundation::llm::ContentBlock::ToolUse { .. }));
353
354        let stop_reason = candidate
355            .finish_reason
356            .as_ref()
357            .map(|r| map_finish_reason(r, has_tool_calls));
358
359        let usage = api_response
360            .usage_metadata
361            .unwrap_or(ApiUsageMetadata {
362                prompt: 0,
363                candidates: 0,
364                cached_content: 0,
365            })
366            .into_usage();
367
368        Ok(ChatOutcome::Success(ChatResponse {
369            id: String::new(),
370            content,
371            model: self.model.clone(),
372            stop_reason,
373            usage,
374        }))
375    }
376
377    fn chat_stream(&self, request: ChatRequest) -> StreamBox<'_> {
378        let served_route = self.route().to_owned();
379        Box::pin(async_stream::stream! {
380            let thinking = match self.resolve_thinking_config(request.thinking.as_ref()) {
381                Ok(thinking) => thinking,
382                Err(error) => {
383                    yield Ok(StreamDelta::Error {
384                        message: error.to_string(),
385                        kind: StreamErrorKind::InvalidRequest,
386                    });
387                    return;
388                }
389            };
390            if let Err(error) = validate_request_attachments(self.provider(), self.model(), &request) {
391                yield Ok(StreamDelta::Error {
392                    message: error.to_string(),
393                    kind: StreamErrorKind::InvalidRequest,
394                });
395                return;
396            }
397            let contents = build_api_contents(&request.messages);
398            let tools = request
399            .tools
400            .as_ref()
401            .map(|t| convert_tools_to_config(t.clone()));
402            let tool_config = request
403                .tool_choice
404                .as_ref()
405                .map(ApiFunctionCallingConfig::from_tool_choice);
406            let system_instruction = if request.system.is_empty() {
407                None
408            } else {
409                Some(ApiContent {
410                    role: None,
411                    parts: vec![ApiPart::Text {
412                        text: request.system.clone(),
413                        thought_signature: None,
414                    }],
415                })
416            };
417
418            let thinking_config = thinking.as_ref().map(map_thinking_config);
419            let (response_mime_type, response_schema) = request
420                .response_format
421                .as_ref()
422                .map_or((None, None), |rf| {
423                    (
424                        Some("application/json"),
425                        Some(gemini_response_schema(&rf.schema)),
426                    )
427                });
428
429            let max_tokens = self.effective_max_tokens(&request);
430            let api_request = ApiGenerateContentRequest {
431                contents: &contents,
432                system_instruction: system_instruction.as_ref(),
433                tools: tools.as_ref().map(std::slice::from_ref),
434                tool_config,
435                generation_config: Some(ApiGenerationConfig {
436                    max_output_tokens: Some(max_tokens),
437                    thinking_config,
438                    response_mime_type,
439                    response_schema,
440                }),
441                cached_content: request.cached_content.as_deref(),
442            };
443
444            log::debug!(
445                "Gemini streaming LLM request model={} max_tokens={}",
446                self.model,
447                max_tokens
448            );
449
450            let stream_builder = self
451                .client
452                .post(format!(
453                    "{}/models/{}:streamGenerateContent",
454                    self.base_url, self.model
455                ))
456                .header("Content-Type", "application/json")
457                .query(&[("alt", "sse")]);
458            let response = match self
459                .apply_auth(stream_builder)
460                .json(&api_request)
461                .send()
462                .await
463            {
464                Ok(r) => r,
465                Err(error) => {
466                    yield Ok(reqwest_error_delta("request failed", &error));
467                    return;
468                }
469            };
470
471            let status = response.status();
472            if !status.is_success() {
473                // Headers are read before the body: `text()` consumes the response.
474                let header_hint = crate::http::retry_after_from_headers(response.headers());
475                let body = response.text().await.unwrap_or_default();
476                let kind = if status == StatusCode::TOO_MANY_REQUESTS {
477                    StreamErrorKind::RateLimited(
478                        header_hint.or_else(|| crate::retry_hints::google_retry_delay(&body)),
479                    )
480                } else if status.is_server_error() {
481                    StreamErrorKind::ServerError
482                } else {
483                    StreamErrorKind::InvalidRequest
484                };
485                log::warn!("Gemini error status={status} body={body}");
486                yield Ok(StreamDelta::Error {
487                    message: body,
488                    kind,
489                });
490                return;
491            }
492
493            let mut inner = data::stream_gemini_response(response);
494            while let Some(item) = futures::StreamExt::next(&mut inner).await {
495                yield match item {
496                    Ok(StreamDelta::Done { stop_reason, .. }) => Ok(StreamDelta::Done {
497                        stop_reason,
498                        served_route: Some(served_route.clone()),
499                    }),
500                    other => other,
501                };
502            }
503        })
504    }
505
506    async fn list_models(&self) -> Result<Vec<crate::provider::ModelInfo>> {
507        // The endpoint paginates (default `pageSize=50`). Request the max page
508        // size and follow `nextPageToken` until exhausted, collecting *raw*
509        // rows. The `generateContent` filter is applied only after every page is
510        // in hand, so server-side truncation cannot hide a chat-capable model.
511        let mut rows: Vec<GeminiModelRow> = Vec::new();
512        let mut page_token: Option<String> = None;
513        for _ in 0..MODELS_MAX_PAGES {
514            let mut query: Vec<(&str, String)> = vec![("pageSize", MODELS_PAGE_SIZE.to_string())];
515            if let Some(token) = &page_token {
516                query.push(("pageToken", token.clone()));
517            }
518            let builder = self
519                .client
520                .get(format!("{}/models", self.base_url))
521                .header("Content-Type", "application/json")
522                .query(&query);
523            let builder = self.apply_auth(builder);
524            let body =
525                crate::impls::model_listing::fetch_model_list_body(builder, "Gemini").await?;
526            let page = parse_models_page(&body)?;
527            rows.extend(page.models);
528            match page.next_page_token {
529                Some(token) if !token.is_empty() => page_token = Some(token),
530                _ => break,
531            }
532        }
533        Ok(finalize_gemini_models(rows))
534    }
535
536    async fn probe_connectivity(&self) -> bool {
537        crate::provider::probe_http_reachability(&self.client, &self.base_url).await
538    }
539
540    fn model(&self) -> &str {
541        &self.model
542    }
543
544    fn provider(&self) -> &'static str {
545        "gemini"
546    }
547
548    fn configured_thinking(&self) -> Option<&ThinkingConfig> {
549        self.thinking.as_ref()
550    }
551}
552
553/// A raw Gemini model row, kept un-filtered so the `generateContent` filter can
554/// be applied only *after* every page has been collected (so server-side page
555/// truncation cannot hide a chat-capable model behind a page boundary).
556#[derive(serde::Deserialize)]
557struct GeminiModelRow {
558    name: String,
559    #[serde(rename = "displayName", default)]
560    display_name: Option<String>,
561    #[serde(rename = "inputTokenLimit", default)]
562    input_token_limit: Option<u32>,
563    #[serde(rename = "outputTokenLimit", default)]
564    output_token_limit: Option<u32>,
565    #[serde(rename = "supportedGenerationMethods", default)]
566    supported_generation_methods: Vec<String>,
567}
568
569/// One page of the Gemini `ListModels` response: raw rows plus the cursor used
570/// to follow pagination.
571struct GeminiModelsPage {
572    models: Vec<GeminiModelRow>,
573    next_page_token: Option<String>,
574}
575
576/// Parse one page of the Gemini `GET /v1beta/models` response body.
577///
578/// The endpoint returns `{ "models": [{ "name": "models/<id>", "displayName",
579/// "inputTokenLimit", "outputTokenLimit", "supportedGenerationMethods" }],
580/// "nextPageToken": "..." }`. It paginates with a default `pageSize` of 50;
581/// `nextPageToken` drives the next request. Raw rows are returned un-filtered so
582/// the caller can apply the `generateContent` filter once all pages are in hand.
583fn parse_models_page(body: &str) -> Result<GeminiModelsPage> {
584    #[derive(serde::Deserialize)]
585    struct ListResponse {
586        #[serde(default)]
587        models: Vec<GeminiModelRow>,
588        #[serde(rename = "nextPageToken", default)]
589        next_page_token: Option<String>,
590    }
591    let parsed: ListResponse = serde_json::from_str(body)
592        .map_err(|e| anyhow::anyhow!("failed to parse Gemini models list: {e}"))?;
593    Ok(GeminiModelsPage {
594        models: parsed.models,
595        next_page_token: parsed.next_page_token,
596    })
597}
598
599/// Filter accumulated rows to chat-capable models and project them into
600/// [`ModelInfo`].
601///
602/// Entries that do not support `generateContent` (e.g. embedding-only models)
603/// are dropped, and the `models/` prefix is stripped from `name` to recover the
604/// bare model id the chat endpoint expects. Applied *after* all pages are
605/// collected so a chat-capable model never gets hidden by page truncation.
606fn finalize_gemini_models(rows: Vec<GeminiModelRow>) -> Vec<crate::provider::ModelInfo> {
607    rows.into_iter()
608        .filter(|row| {
609            row.supported_generation_methods.is_empty()
610                || row
611                    .supported_generation_methods
612                    .iter()
613                    .any(|m| m == "generateContent")
614        })
615        .map(|row| crate::provider::ModelInfo {
616            id: match row.name.strip_prefix("models/") {
617                Some(stripped) => stripped.to_owned(),
618                None => row.name.clone(),
619            },
620            display_name: row.display_name,
621            context_window: row.input_token_limit,
622            max_output_tokens: row.output_token_limit,
623        })
624        .collect()
625}
626
627#[cfg(test)]
628mod tests {
629    use super::*;
630
631    const GEMINI_MODELS_FIXTURE: &str = r#"{
632      "models": [
633        {
634          "name": "models/gemini-2.5-pro",
635          "displayName": "Gemini 2.5 Pro",
636          "inputTokenLimit": 1048576,
637          "outputTokenLimit": 65536,
638          "supportedGenerationMethods": ["generateContent", "countTokens"]
639        },
640        {
641          "name": "models/text-embedding-004",
642          "displayName": "Text Embedding 004",
643          "inputTokenLimit": 2048,
644          "outputTokenLimit": 1,
645          "supportedGenerationMethods": ["embedContent"]
646        }
647      ]
648    }"#;
649
650    #[test]
651    fn parse_models_page_strips_prefix_and_maps_limits() -> anyhow::Result<()> {
652        let page = parse_models_page(GEMINI_MODELS_FIXTURE)?;
653        let models = finalize_gemini_models(page.models);
654        // The embedding-only model is filtered out (no `generateContent`).
655        assert_eq!(models.len(), 1);
656        let pro = &models[0];
657        assert_eq!(pro.id, "gemini-2.5-pro");
658        assert_eq!(pro.display_name.as_deref(), Some("Gemini 2.5 Pro"));
659        assert_eq!(pro.context_window, Some(1_048_576));
660        assert_eq!(pro.max_output_tokens, Some(65_536));
661        assert_eq!(page.next_page_token, None);
662        Ok(())
663    }
664
665    #[tokio::test]
666    async fn list_models_follows_pagination_and_filters_after_all_pages() -> anyhow::Result<()> {
667        use wiremock::matchers::{method, path, query_param, query_param_is_missing};
668        use wiremock::{Mock, MockServer, ResponseTemplate};
669
670        let server = MockServer::start().await;
671
672        // Page 1: a chat model plus an embedding-only model, then a page token.
673        // The embedding model must NOT be filtered out mid-pagination — the
674        // filter runs only after every page is collected.
675        Mock::given(method("GET"))
676            .and(path("/models"))
677            .and(query_param_is_missing("pageToken"))
678            .respond_with(ResponseTemplate::new(200).set_body_string(
679                r#"{
680                  "models": [
681                    {
682                      "name": "models/gemini-2.5-pro",
683                      "displayName": "Gemini 2.5 Pro",
684                      "inputTokenLimit": 1048576,
685                      "outputTokenLimit": 65536,
686                      "supportedGenerationMethods": ["generateContent"]
687                    },
688                    {
689                      "name": "models/text-embedding-004",
690                      "displayName": "Embedding",
691                      "supportedGenerationMethods": ["embedContent"]
692                    }
693                  ],
694                  "nextPageToken": "page-2"
695                }"#,
696            ))
697            .mount(&server)
698            .await;
699
700        // Page 2: requested with `pageToken=page-2`; final page (no token).
701        Mock::given(method("GET"))
702            .and(path("/models"))
703            .and(query_param("pageToken", "page-2"))
704            .respond_with(ResponseTemplate::new(200).set_body_string(
705                r#"{
706                  "models": [
707                    {
708                      "name": "models/gemini-3-flash",
709                      "displayName": "Gemini 3 Flash",
710                      "inputTokenLimit": 1048576,
711                      "outputTokenLimit": 65536,
712                      "supportedGenerationMethods": ["generateContent"]
713                    }
714                  ]
715                }"#,
716            ))
717            .mount(&server)
718            .await;
719
720        let provider = GeminiProvider::new("test-key".to_string(), "gemini-test".to_string())
721            .with_base_url(server.uri());
722        let models = provider.list_models().await?;
723
724        // Both chat models from both pages are returned; the embedding-only
725        // model is dropped by the post-pagination filter.
726        let ids: Vec<&str> = models.iter().map(|m| m.id.as_str()).collect();
727        assert_eq!(ids, vec!["gemini-2.5-pro", "gemini-3-flash"]);
728        Ok(())
729    }
730
731    #[test]
732    fn test_new_creates_provider_with_custom_model() {
733        let provider = GeminiProvider::new("test-api-key".to_string(), "custom-model".to_string());
734
735        assert_eq!(provider.model(), "custom-model");
736        assert_eq!(provider.provider(), "gemini");
737    }
738
739    #[test]
740    fn test_flash_factory_creates_flash_provider() {
741        let provider = GeminiProvider::flash("test-api-key".to_string());
742
743        assert_eq!(provider.model(), MODEL_GEMINI_3_FLASH);
744        assert_eq!(provider.provider(), "gemini");
745    }
746
747    #[test]
748    fn test_flash_lite_factory_creates_flash_lite_provider() {
749        let provider = GeminiProvider::flash_lite("test-api-key".to_string());
750
751        assert_eq!(provider.model(), MODEL_GEMINI_2_FLASH_LITE);
752        assert_eq!(provider.provider(), "gemini");
753    }
754
755    #[test]
756    fn test_flash_lite_31_factory_creates_flash_lite_provider() {
757        let provider = GeminiProvider::flash_lite_31("test-api-key".to_string());
758
759        assert_eq!(provider.model(), MODEL_GEMINI_31_FLASH_LITE);
760        assert_eq!(provider.provider(), "gemini");
761    }
762
763    #[test]
764    fn test_pro_factory_creates_pro_provider() {
765        let provider = GeminiProvider::pro("test-api-key".to_string());
766
767        assert_eq!(provider.model(), MODEL_GEMINI_31_PRO);
768        assert_eq!(provider.provider(), "gemini");
769    }
770
771    #[test]
772    fn test_pro_31_factory_creates_pro_provider() {
773        let provider = GeminiProvider::pro_31("test-api-key".to_string());
774
775        assert_eq!(provider.model(), MODEL_GEMINI_31_PRO);
776        assert_eq!(provider.provider(), "gemini");
777    }
778
779    #[test]
780    fn test_model_constants_have_expected_values() {
781        assert_eq!(MODEL_GEMINI_31_PRO, "gemini-3.1-pro-preview");
782        assert_eq!(MODEL_GEMINI_31_FLASH_LITE, "gemini-3.1-flash-lite-preview");
783        assert_eq!(MODEL_GEMINI_3_FLASH, "gemini-3-flash-preview");
784        assert_eq!(MODEL_GEMINI_3_PRO, "gemini-3.0-pro");
785        assert_eq!(MODEL_GEMINI_25_FLASH, "gemini-2.5-flash");
786        assert_eq!(MODEL_GEMINI_25_PRO, "gemini-2.5-pro");
787        assert_eq!(MODEL_GEMINI_2_FLASH, "gemini-2.0-flash");
788        assert_eq!(MODEL_GEMINI_2_FLASH_LITE, "gemini-2.0-flash-lite");
789    }
790
791    #[test]
792    fn test_gemini_20_models_reject_thinking() {
793        let provider = GeminiProvider::flash_lite("test-api-key".to_string());
794        let error = provider
795            .validate_thinking_config(Some(&ThinkingConfig::new(10_000)))
796            .unwrap_err();
797        assert!(error.to_string().contains("thinking is not supported"));
798    }
799
800    #[test]
801    fn test_default_uses_header_auth() {
802        let provider = GeminiProvider::new("test-key".to_string(), "model".to_string());
803        assert!(
804            provider.use_header_auth,
805            "Default should use header auth for security"
806        );
807    }
808
809    #[test]
810    fn test_provider_is_cloneable() {
811        let provider = GeminiProvider::new("test-api-key".to_string(), "test-model".to_string());
812        let cloned = provider.clone();
813
814        assert_eq!(provider.model(), cloned.model());
815        assert_eq!(provider.provider(), cloned.provider());
816    }
817
818    fn request_with_max_tokens(max_tokens: u32, explicit: bool) -> ChatRequest {
819        ChatRequest {
820            system: String::new(),
821            messages: vec![agent_sdk_foundation::llm::Message::user("hi")],
822            tools: None,
823            max_tokens,
824            max_tokens_explicit: explicit,
825            session_id: None,
826            cached_content: None,
827            thinking: None,
828            tool_choice: None,
829            response_format: None,
830            cache: None,
831        }
832    }
833
834    #[test]
835    fn test_effective_max_tokens_honors_explicit_budget() {
836        let provider = GeminiProvider::pro("test-api-key".to_string());
837        let request = request_with_max_tokens(123, true);
838        assert_eq!(provider.effective_max_tokens(&request), 123);
839    }
840
841    #[test]
842    fn test_effective_max_tokens_uses_default_when_implicit() {
843        // An implicit budget must fall back to the provider/model default, not
844        // be silently capped at ChatRequest::DEFAULT_MAX_TOKENS.
845        let provider = GeminiProvider::pro("test-api-key".to_string());
846        let request = request_with_max_tokens(4096, false);
847        assert_eq!(
848            provider.effective_max_tokens(&request),
849            provider.default_max_tokens()
850        );
851    }
852}