Skip to main content

usage_monitor_cli/provider/
gemini.rs

1//! Gemini (Google AI / Code Assist) usage provider.
2//!
3//! Ports CodexBar's OAuth flow for the Linux CLI: it reads the gemini-cli OAuth
4//! credentials from `~/.gemini/oauth_creds.json` (or an explicit `access_token`
5//! in config), refreshes the access token when expired using the public
6//! gemini-cli OAuth client, then calls Code Assist's `loadCodeAssist` and
7//! `retrieveUserQuota` endpoints to read per-model daily quotas.
8
9use async_trait::async_trait;
10use chrono::{DateTime, Utc};
11
12use crate::error::SpendPanelError;
13use crate::model::{PlanInfo, RateWindow, UsageSnapshot};
14use crate::provider::{ProviderContext, ProviderMetadata, UsageProvider};
15
16/// Public gemini-cli OAuth client (shipped in the open-source `@google/gemini-cli`).
17const GEMINI_CLI_CLIENT_ID: &str =
18    "681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com";
19const GEMINI_CLI_CLIENT_SECRET: &str = "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";
20
21const CLOUDCODE_BASE: &str = "https://cloudcode-pa.googleapis.com";
22const TOKEN_URL: &str = "https://oauth2.googleapis.com/token";
23
24#[derive(Debug, serde::Deserialize)]
25struct QuotaResponse {
26    #[serde(default)]
27    buckets: Option<Vec<QuotaBucket>>,
28}
29
30#[derive(Debug, serde::Deserialize)]
31struct QuotaBucket {
32    #[serde(default, rename = "remainingFraction")]
33    remaining_fraction: Option<f64>,
34    #[serde(default, rename = "resetTime")]
35    reset_time: Option<String>,
36    #[serde(default, rename = "modelId")]
37    model_id: Option<String>,
38}
39
40/// One model's resolved daily quota.
41#[derive(Debug, Clone, PartialEq)]
42struct ModelQuota {
43    model_id: String,
44    percent_left: f64,
45    reset_time: Option<DateTime<Utc>>,
46}
47
48#[derive(Debug, serde::Deserialize)]
49struct OAuthCreds {
50    #[serde(default)]
51    access_token: Option<String>,
52    #[serde(default)]
53    refresh_token: Option<String>,
54    /// Unix-millis expiry, as gemini-cli stores it.
55    #[serde(default)]
56    expiry_date: Option<f64>,
57}
58
59#[derive(Debug, serde::Deserialize)]
60struct RefreshResponse {
61    access_token: String,
62}
63
64fn is_flash_lite(id: &str) -> bool {
65    id.contains("flash-lite")
66}
67fn is_flash(id: &str) -> bool {
68    id.contains("flash") && !is_flash_lite(id)
69}
70fn is_pro(id: &str) -> bool {
71    id.contains("pro")
72}
73
74/// Gemini Code Assist usage provider.
75pub struct GeminiProvider {
76    metadata: ProviderMetadata,
77    /// Base for cloudcode-pa endpoints (overridable in tests).
78    cloudcode_base: Option<String>,
79    /// Base for the OAuth token endpoint (overridable in tests).
80    token_url: Option<String>,
81}
82
83impl GeminiProvider {
84    pub fn new() -> Self {
85        Self {
86            metadata: ProviderMetadata {
87                id: "gemini",
88                name: "Google Gemini",
89                description: "Gemini Code Assist daily quota monitor (gemini-cli OAuth)",
90                auth_methods: &["oauth", "access_token", "env"],
91                website: Some("https://aistudio.google.com"),
92            },
93            cloudcode_base: None,
94            token_url: None,
95        }
96    }
97
98    /// Points cloudcode + token endpoints at a test server.
99    pub fn with_base_url(url: &str) -> Self {
100        let mut p = Self::new();
101        p.cloudcode_base = Some(url.to_string());
102        p.token_url = Some(format!("{}/token", url.trim_end_matches('/')));
103        p
104    }
105
106    fn cloudcode_base(&self) -> &str {
107        self.cloudcode_base.as_deref().unwrap_or(CLOUDCODE_BASE)
108    }
109
110    fn token_url(&self) -> &str {
111        self.token_url.as_deref().unwrap_or(TOKEN_URL)
112    }
113
114    fn creds_path(ctx: &ProviderContext) -> std::path::PathBuf {
115        if let Some(p) = ctx.config.get("credentials_path").filter(|v| !v.is_empty()) {
116            return std::path::PathBuf::from(p);
117        }
118        let home = std::env::var("HOME").unwrap_or_default();
119        std::path::Path::new(&home).join(".gemini/oauth_creds.json")
120    }
121
122    fn build_client(ctx: &ProviderContext) -> Result<reqwest::Client, SpendPanelError> {
123        reqwest::Client::builder()
124            .timeout(std::time::Duration::from_secs(ctx.timeout_secs))
125            .build()
126            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))
127    }
128
129    /// Resolves a usable access token: explicit config token, or the creds file
130    /// (refreshing when expired).
131    async fn resolve_access_token(
132        &self,
133        ctx: &ProviderContext,
134        client: &reqwest::Client,
135    ) -> Result<String, SpendPanelError> {
136        for key in ["access_token", "token"] {
137            if let Some(value) = ctx
138                .config
139                .get(key)
140                .map(|s| s.trim())
141                .filter(|s| !s.is_empty())
142            {
143                return Ok(value.to_string());
144            }
145        }
146
147        let path = Self::creds_path(ctx);
148        let data = std::fs::read_to_string(&path).map_err(|_| {
149            SpendPanelError::AuthFailed(
150                "gemini".into(),
151                format!(
152                    "no access_token in config and no credentials at {}",
153                    path.display()
154                ),
155            )
156        })?;
157        let creds: OAuthCreds = serde_json::from_str(&data)
158            .map_err(|e| SpendPanelError::ParseError("gemini".into(), e.to_string()))?;
159
160        let expired = creds
161            .expiry_date
162            .map(|ms| (ms / 1000.0) < Utc::now().timestamp() as f64)
163            .unwrap_or(true);
164        let token = creds.access_token.clone().filter(|t| !t.is_empty());
165
166        if let Some(token) = token.filter(|_| !expired) {
167            return Ok(token);
168        }
169
170        let refresh = creds
171            .refresh_token
172            .filter(|t| !t.is_empty())
173            .ok_or_else(|| {
174                SpendPanelError::AuthFailed(
175                    "gemini".into(),
176                    "access token expired and no refresh_token available; re-run gemini login"
177                        .into(),
178                )
179            })?;
180        self.refresh_access_token(client, &refresh).await
181    }
182
183    async fn refresh_access_token(
184        &self,
185        client: &reqwest::Client,
186        refresh_token: &str,
187    ) -> Result<String, SpendPanelError> {
188        let params = [
189            ("client_id", GEMINI_CLI_CLIENT_ID),
190            ("client_secret", GEMINI_CLI_CLIENT_SECRET),
191            ("refresh_token", refresh_token),
192            ("grant_type", "refresh_token"),
193        ];
194        let resp = client
195            .post(self.token_url())
196            .form(&params)
197            .send()
198            .await
199            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
200        let status = resp.status();
201        let body = resp
202            .text()
203            .await
204            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
205        if !status.is_success() {
206            return Err(SpendPanelError::AuthFailed(
207                "gemini".into(),
208                format!("token refresh failed (HTTP {})", status.as_u16()),
209            ));
210        }
211        let parsed: RefreshResponse = serde_json::from_str(&body)
212            .map_err(|e| SpendPanelError::ParseError("gemini".into(), e.to_string()))?;
213        Ok(parsed.access_token)
214    }
215
216    /// Loads the Code Assist project id (best-effort; `None` on any failure).
217    async fn load_project_id(
218        &self,
219        client: &reqwest::Client,
220        access_token: &str,
221    ) -> Option<String> {
222        let url = format!(
223            "{}/v1internal:loadCodeAssist",
224            self.cloudcode_base().trim_end_matches('/')
225        );
226        let resp = client
227            .post(url)
228            .header("Authorization", format!("Bearer {}", access_token))
229            .header("Content-Type", "application/json")
230            .body(r#"{"metadata":{"ideType":"GEMINI_CLI","pluginType":"GEMINI"}}"#)
231            .send()
232            .await
233            .ok()?;
234        if !resp.status().is_success() {
235            return None;
236        }
237        let json: serde_json::Value = resp.json().await.ok()?;
238        let project = json.get("cloudaicompanionProject");
239        match project {
240            Some(serde_json::Value::String(s)) if !s.trim().is_empty() => {
241                Some(s.trim().to_string())
242            }
243            Some(serde_json::Value::Object(o)) => o
244                .get("id")
245                .or_else(|| o.get("projectId"))
246                .and_then(|v| v.as_str())
247                .filter(|s| !s.trim().is_empty())
248                .map(|s| s.trim().to_string()),
249            _ => None,
250        }
251    }
252
253    async fn retrieve_quota(
254        &self,
255        client: &reqwest::Client,
256        access_token: &str,
257        project_id: Option<&str>,
258    ) -> Result<QuotaResponse, SpendPanelError> {
259        let url = format!(
260            "{}/v1internal:retrieveUserQuota",
261            self.cloudcode_base().trim_end_matches('/')
262        );
263        let body = match project_id {
264            Some(id) => format!(r#"{{"project": "{}"}}"#, id),
265            None => "{}".to_string(),
266        };
267        let resp = client
268            .post(url)
269            .header("Authorization", format!("Bearer {}", access_token))
270            .header("Content-Type", "application/json")
271            .body(body)
272            .send()
273            .await
274            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
275        let status = resp.status();
276        let text = resp
277            .text()
278            .await
279            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
280        if status == reqwest::StatusCode::UNAUTHORIZED {
281            return Err(SpendPanelError::AuthFailed(
282                "gemini".into(),
283                "access token rejected (HTTP 401)".into(),
284            ));
285        }
286        if !status.is_success() {
287            return Err(SpendPanelError::ProviderError(
288                "gemini".into(),
289                format!("HTTP {}: {}", status, text),
290            ));
291        }
292        serde_json::from_str(&text)
293            .map_err(|e| SpendPanelError::ParseError("gemini".into(), e.to_string()))
294    }
295
296    /// Groups buckets by model (keeping the lowest remaining fraction per model).
297    fn parse_quota(resp: &QuotaResponse) -> Result<Vec<ModelQuota>, SpendPanelError> {
298        let buckets = resp
299            .buckets
300            .as_deref()
301            .filter(|b| !b.is_empty())
302            .ok_or_else(|| {
303                SpendPanelError::ParseError("gemini".into(), "no quota buckets in response".into())
304            })?;
305
306        let mut map: std::collections::BTreeMap<String, (f64, Option<String>)> =
307            std::collections::BTreeMap::new();
308        for bucket in buckets {
309            let (Some(model_id), Some(fraction)) = (&bucket.model_id, bucket.remaining_fraction)
310            else {
311                continue;
312            };
313            map.entry(model_id.clone())
314                .and_modify(|existing| {
315                    if fraction < existing.0 {
316                        *existing = (fraction, bucket.reset_time.clone());
317                    }
318                })
319                .or_insert((fraction, bucket.reset_time.clone()));
320        }
321
322        Ok(map
323            .into_iter()
324            .map(|(model_id, (fraction, reset))| ModelQuota {
325                model_id,
326                percent_left: fraction * 100.0,
327                reset_time: reset
328                    .as_deref()
329                    .and_then(|s| DateTime::parse_from_rfc3339(s).ok())
330                    .map(|d| d.with_timezone(&Utc)),
331            })
332            .collect())
333    }
334
335    fn snapshot_from_quotas(quotas: &[ModelQuota]) -> UsageSnapshot {
336        let lowest = |pred: fn(&str) -> bool| -> Option<&ModelQuota> {
337            quotas
338                .iter()
339                .filter(|q| pred(&q.model_id.to_lowercase()))
340                .min_by(|a, b| a.percent_left.total_cmp(&b.percent_left))
341        };
342        let window = |q: &ModelQuota, label: &str| -> RateWindow {
343            let used = (100.0 - q.percent_left).clamp(0.0, 100.0).round() as u64;
344            let mut w = RateWindow::new(used, 100, label.to_string(), 1440);
345            w.resets_at = q.reset_time;
346            w
347        };
348
349        let mut snapshot = UsageSnapshot::new("gemini");
350        if let Some(pro) = lowest(is_pro) {
351            snapshot.primary_rate_window = Some(window(pro, "Gemini Pro"));
352        }
353        if let Some(flash) = lowest(is_flash) {
354            snapshot.secondary_rate_window = Some(window(flash, "Gemini Flash"));
355        }
356        if let Some(lite) = lowest(is_flash_lite) {
357            snapshot.tertiary_rate_window = Some(window(lite, "Gemini Flash Lite"));
358        }
359        // Fall back to a plain window when no model matched the known families.
360        let needs_fallback = snapshot.primary_rate_window.is_none();
361        if let Some(any) = quotas
362            .iter()
363            .min_by(|a, b| a.percent_left.total_cmp(&b.percent_left))
364            .filter(|_| needs_fallback)
365        {
366            let label = any.model_id.clone();
367            snapshot.primary_rate_window = Some(window(any, &label));
368        }
369        snapshot
370    }
371}
372
373impl Default for GeminiProvider {
374    fn default() -> Self {
375        Self::new()
376    }
377}
378
379#[async_trait]
380impl UsageProvider for GeminiProvider {
381    fn metadata(&self) -> &ProviderMetadata {
382        &self.metadata
383    }
384
385    fn detect_credentials(&self) -> bool {
386        let home = std::env::var("HOME").unwrap_or_default();
387        std::path::Path::new(&home)
388            .join(".gemini/oauth_creds.json")
389            .exists()
390    }
391
392    async fn fetch_usage(&self, ctx: &ProviderContext) -> Result<UsageSnapshot, SpendPanelError> {
393        let client = Self::build_client(ctx)?;
394        let access_token = self.resolve_access_token(ctx, &client).await?;
395
396        let project_id = match ctx.config.get("project").filter(|v| !v.is_empty()) {
397            Some(p) => Some(p.clone()),
398            None => self.load_project_id(&client, &access_token).await,
399        };
400
401        let quota = self
402            .retrieve_quota(&client, &access_token, project_id.as_deref())
403            .await?;
404        let quotas = Self::parse_quota(&quota)?;
405        let mut snapshot = Self::snapshot_from_quotas(&quotas);
406        snapshot.plan = Some(PlanInfo {
407            name: "Code Assist".into(),
408            tier: None,
409            features: Vec::new(),
410            price: None,
411            currency: None,
412            billing_period: Some("daily".into()),
413        });
414        Ok(snapshot)
415    }
416}
417
418#[cfg(test)]
419mod tests {
420    use super::*;
421    use pretty_assertions::assert_eq;
422    use wiremock::matchers::{method, path};
423    use wiremock::{Mock, MockServer, ResponseTemplate};
424
425    const QUOTA: &str = r#"{
426      "buckets": [
427        {"modelId": "gemini-2.5-pro", "remainingFraction": 0.4, "resetTime": "2026-06-14T00:00:00Z"},
428        {"modelId": "gemini-2.5-pro", "remainingFraction": 0.2, "resetTime": "2026-06-14T00:00:00Z"},
429        {"modelId": "gemini-2.5-flash", "remainingFraction": 0.9, "resetTime": "2026-06-14T00:00:00Z"},
430        {"modelId": "gemini-2.5-flash-lite", "remainingFraction": 1.0, "resetTime": "2026-06-14T00:00:00Z"}
431      ]
432    }"#;
433
434    fn quota(body: &str) -> QuotaResponse {
435        serde_json::from_str(body).unwrap()
436    }
437
438    #[test]
439    fn test_metadata() {
440        let p = GeminiProvider::new();
441        assert_eq!(p.metadata().id, "gemini");
442    }
443
444    #[test]
445    fn test_model_classifiers() {
446        assert!(is_flash_lite("gemini-2.5-flash-lite"));
447        assert!(is_flash("gemini-2.5-flash"));
448        assert!(!is_flash("gemini-2.5-flash-lite"));
449        assert!(is_pro("gemini-2.5-pro"));
450    }
451
452    #[test]
453    fn test_parse_quota_keeps_lowest_per_model() {
454        let quotas = GeminiProvider::parse_quota(&quota(QUOTA)).unwrap();
455        let pro = quotas.iter().find(|q| q.model_id.contains("pro")).unwrap();
456        assert_eq!(pro.percent_left, 20.0); // lowest of 0.4/0.2
457    }
458
459    #[test]
460    fn test_parse_quota_empty_is_error() {
461        let err = GeminiProvider::parse_quota(&quota(r#"{"buckets":[]}"#)).unwrap_err();
462        assert!(matches!(err, SpendPanelError::ParseError(_, _)));
463    }
464
465    #[test]
466    fn test_snapshot_maps_families() {
467        let quotas = GeminiProvider::parse_quota(&quota(QUOTA)).unwrap();
468        let snapshot = GeminiProvider::snapshot_from_quotas(&quotas);
469        // pro 20% left → 80% used
470        assert_eq!(snapshot.primary_rate_window.unwrap().used, Some(80));
471        // flash 90% left → 10% used
472        assert_eq!(snapshot.secondary_rate_window.unwrap().used, Some(10));
473        // flash-lite 100% left → 0% used
474        assert_eq!(snapshot.tertiary_rate_window.unwrap().used, Some(0));
475    }
476
477    #[test]
478    fn test_snapshot_unknown_family_falls_back_to_primary() {
479        // A model that matches no known family still populates the primary lane.
480        let quotas = GeminiProvider::parse_quota(&quota(
481            r#"{"buckets":[{"modelId":"some-experimental-model","remainingFraction":0.3}]}"#,
482        ))
483        .unwrap();
484        let snapshot = GeminiProvider::snapshot_from_quotas(&quotas);
485        assert_eq!(snapshot.primary_rate_window.unwrap().used, Some(70));
486    }
487
488    #[tokio::test]
489    async fn test_fetch_usage_with_config_token() {
490        let server = MockServer::start().await;
491        Mock::given(method("POST"))
492            .and(path("/v1internal:loadCodeAssist"))
493            .respond_with(ResponseTemplate::new(200).set_body_raw(
494                r#"{"cloudaicompanionProject":"proj-1"}"#,
495                "application/json",
496            ))
497            .mount(&server)
498            .await;
499        Mock::given(method("POST"))
500            .and(path("/v1internal:retrieveUserQuota"))
501            .respond_with(ResponseTemplate::new(200).set_body_raw(QUOTA, "application/json"))
502            .mount(&server)
503            .await;
504
505        let provider = GeminiProvider::with_base_url(&server.uri());
506        let mut ctx = ProviderContext::new();
507        ctx.config.insert("access_token".into(), "ya29-test".into());
508        let snapshot = provider.fetch_usage(&ctx).await.unwrap();
509        assert_eq!(snapshot.primary_rate_window.unwrap().used, Some(80));
510    }
511
512    #[tokio::test]
513    async fn test_retrieve_quota_401_is_auth_failed() {
514        let server = MockServer::start().await;
515        Mock::given(method("POST"))
516            .and(path("/v1internal:loadCodeAssist"))
517            .respond_with(ResponseTemplate::new(500))
518            .mount(&server)
519            .await;
520        Mock::given(method("POST"))
521            .and(path("/v1internal:retrieveUserQuota"))
522            .respond_with(ResponseTemplate::new(401))
523            .mount(&server)
524            .await;
525
526        let provider = GeminiProvider::with_base_url(&server.uri());
527        let mut ctx = ProviderContext::new();
528        ctx.config.insert("access_token".into(), "bad".into());
529        let err = provider.fetch_usage(&ctx).await.unwrap_err();
530        assert!(matches!(err, SpendPanelError::AuthFailed(_, _)));
531    }
532}