Skip to main content

usage_monitor_cli/provider/
antigravity.rs

1//! Antigravity (Google Code Assist) usage provider.
2//!
3//! Ports CodexBar's remote-usage flow for Linux: it reads Antigravity's Google
4//! OAuth credentials (default `~/.codexbar/antigravity/oauth_creds.json`, or an
5//! explicit `access_token`), refreshes when expired, then reads per-model daily
6//! quotas from Code Assist's `fetchAvailableModels`, falling back to
7//! `retrieveUserQuota` buckets when models carry no consumed quota.
8
9use async_trait::async_trait;
10use chrono::{DateTime, Utc};
11
12use crate::error::SpendPanelError;
13use crate::model::{NamedRateWindow, PlanInfo, RateWindow, UsageSnapshot};
14use crate::provider::{ProviderContext, ProviderMetadata, UsageProvider};
15
16const CLOUDCODE_BASE: &str = "https://cloudcode-pa.googleapis.com";
17const TOKEN_URL: &str = "https://oauth2.googleapis.com/token";
18
19#[derive(Debug, serde::Deserialize)]
20struct OAuthCreds {
21    #[serde(default, alias = "accessToken")]
22    access_token: Option<String>,
23    #[serde(default, alias = "refreshToken")]
24    refresh_token: Option<String>,
25    #[serde(default, alias = "expiresAt")]
26    expiry_date: Option<f64>,
27    #[serde(default, alias = "projectId", alias = "project_id")]
28    project_id: Option<String>,
29    #[serde(default, alias = "clientId")]
30    client_id: Option<String>,
31    #[serde(default, alias = "clientSecret")]
32    client_secret: Option<String>,
33}
34
35#[derive(Debug, serde::Deserialize)]
36struct RefreshResponse {
37    access_token: String,
38}
39
40#[derive(Debug, serde::Deserialize)]
41struct FetchAvailableModelsResponse {
42    #[serde(default)]
43    models: Option<std::collections::HashMap<String, RemoteModel>>,
44}
45
46#[derive(Debug, serde::Deserialize)]
47struct RemoteModel {
48    #[serde(default, rename = "displayName")]
49    display_name: Option<String>,
50    #[serde(default)]
51    label: Option<String>,
52    #[serde(default, rename = "quotaInfo")]
53    quota_info: Option<RemoteQuotaInfo>,
54}
55
56#[derive(Debug, serde::Deserialize)]
57struct RemoteQuotaInfo {
58    #[serde(default, rename = "remainingFraction")]
59    remaining_fraction: Option<f64>,
60    #[serde(default, rename = "resetTime")]
61    reset_time: Option<String>,
62}
63
64#[derive(Debug, serde::Deserialize)]
65struct RetrieveUserQuotaResponse {
66    #[serde(default)]
67    buckets: Option<Vec<RetrieveUserQuotaBucket>>,
68}
69
70#[derive(Debug, serde::Deserialize)]
71struct RetrieveUserQuotaBucket {
72    #[serde(default, rename = "modelId")]
73    model_id: Option<String>,
74    #[serde(default, rename = "remainingFraction")]
75    remaining_fraction: Option<f64>,
76    #[serde(default, rename = "resetTime")]
77    reset_time: Option<String>,
78}
79
80/// One model's resolved daily quota.
81#[derive(Debug, Clone, PartialEq)]
82struct ModelQuota {
83    model_id: String,
84    label: String,
85    remaining_fraction: Option<f64>,
86    reset_time: Option<DateTime<Utc>>,
87}
88
89impl ModelQuota {
90    fn percent_left(&self) -> f64 {
91        self.remaining_fraction.unwrap_or(1.0) * 100.0
92    }
93}
94
95/// Antigravity Code Assist usage provider.
96pub struct AntigravityProvider {
97    metadata: ProviderMetadata,
98    cloudcode_base: Option<String>,
99    token_url: Option<String>,
100}
101
102impl AntigravityProvider {
103    pub fn new() -> Self {
104        Self {
105            metadata: ProviderMetadata {
106                id: "antigravity",
107                name: "Antigravity",
108                description: "Antigravity Code Assist daily quota monitor (Google OAuth)",
109                auth_methods: &["oauth", "access_token", "env"],
110                website: Some("https://antigravity.google"),
111            },
112            cloudcode_base: None,
113            token_url: None,
114        }
115    }
116
117    pub fn with_base_url(url: &str) -> Self {
118        let mut p = Self::new();
119        p.cloudcode_base = Some(url.to_string());
120        p.token_url = Some(format!("{}/token", url.trim_end_matches('/')));
121        p
122    }
123
124    fn cloudcode_base(&self) -> &str {
125        self.cloudcode_base.as_deref().unwrap_or(CLOUDCODE_BASE)
126    }
127
128    fn token_url(&self) -> &str {
129        self.token_url.as_deref().unwrap_or(TOKEN_URL)
130    }
131
132    fn creds_path(ctx: &ProviderContext) -> std::path::PathBuf {
133        if let Some(p) = ctx.config.get("credentials_path").filter(|v| !v.is_empty()) {
134            return std::path::PathBuf::from(p);
135        }
136        let home = std::env::var("HOME").unwrap_or_default();
137        std::path::Path::new(&home).join(".codexbar/antigravity/oauth_creds.json")
138    }
139
140    fn build_client(ctx: &ProviderContext) -> Result<reqwest::Client, SpendPanelError> {
141        reqwest::Client::builder()
142            .timeout(std::time::Duration::from_secs(ctx.timeout_secs))
143            .build()
144            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))
145    }
146
147    /// (access_token, project_id) resolved from config or the creds file.
148    async fn resolve_auth(
149        &self,
150        ctx: &ProviderContext,
151        client: &reqwest::Client,
152    ) -> Result<(String, Option<String>), SpendPanelError> {
153        for key in ["access_token", "token"] {
154            if let Some(value) = ctx
155                .config
156                .get(key)
157                .map(|s| s.trim())
158                .filter(|s| !s.is_empty())
159            {
160                let project = ctx.config.get("project").filter(|v| !v.is_empty()).cloned();
161                return Ok((value.to_string(), project));
162            }
163        }
164
165        let path = Self::creds_path(ctx);
166        let data = std::fs::read_to_string(&path).map_err(|_| {
167            SpendPanelError::AuthFailed(
168                "antigravity".into(),
169                format!(
170                    "no access_token in config and no credentials at {}",
171                    path.display()
172                ),
173            )
174        })?;
175        let creds: OAuthCreds = serde_json::from_str(&data)
176            .map_err(|e| SpendPanelError::ParseError("antigravity".into(), e.to_string()))?;
177
178        let project = ctx
179            .config
180            .get("project")
181            .filter(|v| !v.is_empty())
182            .cloned()
183            .or_else(|| creds.project_id.clone());
184
185        let expired = creds
186            .expiry_date
187            .map(|ms| (ms / 1000.0) < Utc::now().timestamp() as f64)
188            .unwrap_or(true);
189        if let Some(token) = creds
190            .access_token
191            .clone()
192            .filter(|t| !t.is_empty() && !expired)
193        {
194            return Ok((token, project));
195        }
196
197        let refresh = creds
198            .refresh_token
199            .clone()
200            .filter(|t| !t.is_empty())
201            .ok_or_else(|| {
202                SpendPanelError::AuthFailed(
203                    "antigravity".into(),
204                    "access token expired and no refresh_token available; re-run antigravity login"
205                        .into(),
206                )
207            })?;
208        let token = self
209            .refresh_access_token(ctx, client, &creds, &refresh)
210            .await?;
211        Ok((token, project))
212    }
213
214    async fn refresh_access_token(
215        &self,
216        ctx: &ProviderContext,
217        client: &reqwest::Client,
218        creds: &OAuthCreds,
219        refresh_token: &str,
220    ) -> Result<String, SpendPanelError> {
221        let client_id = ctx
222            .config
223            .get("client_id")
224            .filter(|v| !v.is_empty())
225            .cloned()
226            .or_else(|| {
227                std::env::var("ANTIGRAVITY_OAUTH_CLIENT_ID")
228                    .ok()
229                    .filter(|v| !v.is_empty())
230            })
231            .or_else(|| creds.client_id.clone());
232        let client_secret = ctx
233            .config
234            .get("client_secret")
235            .filter(|v| !v.is_empty())
236            .cloned()
237            .or_else(|| {
238                std::env::var("ANTIGRAVITY_OAUTH_CLIENT_SECRET")
239                    .ok()
240                    .filter(|v| !v.is_empty())
241            })
242            .or_else(|| creds.client_secret.clone());
243
244        let (Some(client_id), Some(client_secret)) = (client_id, client_secret) else {
245            return Err(SpendPanelError::AuthFailed(
246                "antigravity".into(),
247                "OAuth client not configured; set ANTIGRAVITY_OAUTH_CLIENT_ID/SECRET or store them in the credentials".into(),
248            ));
249        };
250
251        let params = [
252            ("client_id", client_id.as_str()),
253            ("client_secret", client_secret.as_str()),
254            ("refresh_token", refresh_token),
255            ("grant_type", "refresh_token"),
256        ];
257        let resp = client
258            .post(self.token_url())
259            .form(&params)
260            .send()
261            .await
262            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
263        let status = resp.status();
264        let body = resp
265            .text()
266            .await
267            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
268        if !status.is_success() {
269            return Err(SpendPanelError::AuthFailed(
270                "antigravity".into(),
271                format!("token refresh failed (HTTP {})", status.as_u16()),
272            ));
273        }
274        let parsed: RefreshResponse = serde_json::from_str(&body)
275            .map_err(|e| SpendPanelError::ParseError("antigravity".into(), e.to_string()))?;
276        Ok(parsed.access_token)
277    }
278
279    async fn post_json(
280        &self,
281        client: &reqwest::Client,
282        endpoint: &str,
283        access_token: &str,
284        body: String,
285    ) -> Result<(reqwest::StatusCode, String), SpendPanelError> {
286        let url = format!(
287            "{}/v1internal:{}",
288            self.cloudcode_base().trim_end_matches('/'),
289            endpoint
290        );
291        let resp = client
292            .post(url)
293            .header("Authorization", format!("Bearer {}", access_token))
294            .header("Content-Type", "application/json")
295            .body(body)
296            .send()
297            .await
298            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
299        let status = resp.status();
300        let text = resp
301            .text()
302            .await
303            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
304        Ok((status, text))
305    }
306
307    fn quota_body(project_id: Option<&str>) -> String {
308        match project_id {
309            Some(id) => format!(r#"{{"project": "{}"}}"#, id),
310            None => "{}".to_string(),
311        }
312    }
313
314    /// Resolves per-model quotas: prefer fetchAvailableModels, fall back to
315    /// retrieveUserQuota buckets when models report no consumed quota.
316    async fn fetch_model_quotas(
317        &self,
318        client: &reqwest::Client,
319        access_token: &str,
320        project_id: Option<&str>,
321    ) -> Result<Vec<ModelQuota>, SpendPanelError> {
322        let body = Self::quota_body(project_id);
323        let (status, text) = self
324            .post_json(client, "fetchAvailableModels", access_token, body.clone())
325            .await?;
326
327        if status == reqwest::StatusCode::UNAUTHORIZED {
328            return Err(SpendPanelError::AuthFailed(
329                "antigravity".into(),
330                "access token rejected (HTTP 401)".into(),
331            ));
332        }
333
334        let from_models = if status.is_success() {
335            serde_json::from_str::<FetchAvailableModelsResponse>(&text)
336                .ok()
337                .map(|r| Self::parse_models(&r))
338                .unwrap_or_default()
339        } else {
340            Vec::new()
341        };
342
343        // When every model is full (or none were returned), the consumed quota
344        // lives in retrieveUserQuota — query it as the authoritative source.
345        let all_full = from_models
346            .iter()
347            .all(|q| q.remaining_fraction.map(|f| f >= 0.999).unwrap_or(true));
348        if from_models.is_empty() || all_full {
349            let (qstatus, qtext) = self
350                .post_json(client, "retrieveUserQuota", access_token, body)
351                .await?;
352            if qstatus == reqwest::StatusCode::UNAUTHORIZED {
353                return Err(SpendPanelError::AuthFailed(
354                    "antigravity".into(),
355                    "access token rejected (HTTP 401)".into(),
356                ));
357            }
358            let parsed = qstatus
359                .is_success()
360                .then(|| serde_json::from_str::<RetrieveUserQuotaResponse>(&qtext).ok())
361                .flatten();
362            if let Some(parsed) = parsed {
363                let buckets = Self::parse_buckets(&parsed);
364                if !buckets.is_empty() {
365                    return Ok(buckets);
366                }
367            }
368            if from_models.is_empty() {
369                return Err(SpendPanelError::ProviderError(
370                    "antigravity".into(),
371                    "no model quotas available (fetchAvailableModels and retrieveUserQuota both empty)".into(),
372                ));
373            }
374        }
375        Ok(from_models)
376    }
377
378    fn parse_models(resp: &FetchAvailableModelsResponse) -> Vec<ModelQuota> {
379        let Some(models) = &resp.models else {
380            return Vec::new();
381        };
382        let mut quotas: Vec<ModelQuota> = models
383            .iter()
384            .filter_map(|(id, model)| {
385                let quota = model.quota_info.as_ref()?;
386                let label = model
387                    .display_name
388                    .as_deref()
389                    .filter(|s| !s.trim().is_empty())
390                    .or(model.label.as_deref().filter(|s| !s.trim().is_empty()))
391                    .unwrap_or(id)
392                    .to_string();
393                Some(ModelQuota {
394                    model_id: id.clone(),
395                    label,
396                    remaining_fraction: quota.remaining_fraction,
397                    reset_time: parse_reset(quota.reset_time.as_deref()),
398                })
399            })
400            .collect();
401        quotas.sort_by(|a, b| a.model_id.cmp(&b.model_id));
402        quotas
403    }
404
405    fn parse_buckets(resp: &RetrieveUserQuotaResponse) -> Vec<ModelQuota> {
406        let Some(buckets) = &resp.buckets else {
407            return Vec::new();
408        };
409        let mut map: std::collections::BTreeMap<String, (Option<f64>, Option<String>)> =
410            std::collections::BTreeMap::new();
411        for bucket in buckets {
412            let Some(model_id) = bucket
413                .model_id
414                .as_deref()
415                .map(str::trim)
416                .filter(|s| !s.is_empty())
417            else {
418                continue;
419            };
420            let next = (bucket.remaining_fraction, bucket.reset_time.clone());
421            map.entry(model_id.to_string())
422                .and_modify(|existing| {
423                    let cur = existing.0.unwrap_or(f64::MAX);
424                    let nv = next.0.unwrap_or(f64::MAX);
425                    if nv < cur {
426                        *existing = next.clone();
427                    }
428                })
429                .or_insert(next);
430        }
431        map.into_iter()
432            .map(|(model_id, (fraction, reset))| ModelQuota {
433                label: model_id.clone(),
434                model_id,
435                remaining_fraction: fraction,
436                reset_time: parse_reset(reset.as_deref()),
437            })
438            .collect()
439    }
440
441    fn snapshot_from_quotas(quotas: &[ModelQuota]) -> UsageSnapshot {
442        let mut sorted: Vec<&ModelQuota> = quotas.iter().collect();
443        sorted.sort_by(|a, b| a.percent_left().total_cmp(&b.percent_left()));
444
445        let window = |q: &ModelQuota| -> RateWindow {
446            let used = (100.0 - q.percent_left()).clamp(0.0, 100.0).round() as u64;
447            let mut w = RateWindow::new(used, 100, q.label.clone(), 1440);
448            w.resets_at = q.reset_time;
449            w
450        };
451
452        let mut snapshot = UsageSnapshot::new("antigravity");
453        let mut iter = sorted.into_iter();
454        if let Some(q) = iter.next() {
455            snapshot.primary_rate_window = Some(window(q));
456        }
457        if let Some(q) = iter.next() {
458            snapshot.secondary_rate_window = Some(window(q));
459        }
460        if let Some(q) = iter.next() {
461            snapshot.tertiary_rate_window = Some(window(q));
462        }
463        for q in iter {
464            snapshot.extra_rate_windows.push(NamedRateWindow {
465                id: q.model_id.clone(),
466                label: q.label.clone(),
467                window: window(q),
468            });
469        }
470        snapshot
471    }
472}
473
474fn parse_reset(s: Option<&str>) -> Option<DateTime<Utc>> {
475    let raw = s?;
476    DateTime::parse_from_rfc3339(raw)
477        .ok()
478        .map(|d| d.with_timezone(&Utc))
479}
480
481impl Default for AntigravityProvider {
482    fn default() -> Self {
483        Self::new()
484    }
485}
486
487#[async_trait]
488impl UsageProvider for AntigravityProvider {
489    fn metadata(&self) -> &ProviderMetadata {
490        &self.metadata
491    }
492
493    fn detect_credentials(&self) -> bool {
494        let home = std::env::var("HOME").unwrap_or_default();
495        std::path::Path::new(&home)
496            .join(".codexbar/antigravity/oauth_creds.json")
497            .exists()
498    }
499
500    async fn fetch_usage(&self, ctx: &ProviderContext) -> Result<UsageSnapshot, SpendPanelError> {
501        let client = Self::build_client(ctx)?;
502        let (access_token, project_id) = self.resolve_auth(ctx, &client).await?;
503        let quotas = self
504            .fetch_model_quotas(&client, &access_token, project_id.as_deref())
505            .await?;
506        let mut snapshot = Self::snapshot_from_quotas(&quotas);
507        snapshot.plan = Some(PlanInfo {
508            name: "Code Assist".into(),
509            tier: None,
510            features: Vec::new(),
511            price: None,
512            currency: None,
513            billing_period: Some("daily".into()),
514        });
515        Ok(snapshot)
516    }
517}
518
519#[cfg(test)]
520mod tests {
521    use super::*;
522    use pretty_assertions::assert_eq;
523    use wiremock::matchers::{method, path};
524    use wiremock::{Mock, MockServer, ResponseTemplate};
525
526    const MODELS: &str = r#"{
527      "models": {
528        "claude-sonnet": {"displayName": "Claude Sonnet", "quotaInfo": {"remainingFraction": 0.25, "resetTime": "2026-06-14T00:00:00Z"}},
529        "gemini-pro": {"displayName": "Gemini Pro", "quotaInfo": {"remainingFraction": 0.8}}
530      }
531    }"#;
532
533    const BUCKETS: &str = r#"{
534      "buckets": [
535        {"modelId": "claude-sonnet", "remainingFraction": 0.5, "resetTime": "2026-06-14T00:00:00Z"},
536        {"modelId": "claude-sonnet", "remainingFraction": 0.3},
537        {"modelId": "gemini-pro", "remainingFraction": 0.6}
538      ]
539    }"#;
540
541    #[test]
542    fn test_metadata() {
543        assert_eq!(AntigravityProvider::new().metadata().id, "antigravity");
544    }
545
546    #[test]
547    fn test_parse_models() {
548        let resp: FetchAvailableModelsResponse = serde_json::from_str(MODELS).unwrap();
549        let quotas = AntigravityProvider::parse_models(&resp);
550        assert_eq!(quotas.len(), 2);
551        let claude = quotas
552            .iter()
553            .find(|q| q.model_id == "claude-sonnet")
554            .unwrap();
555        assert_eq!(claude.label, "Claude Sonnet");
556        assert_eq!(claude.percent_left(), 25.0);
557    }
558
559    #[test]
560    fn test_parse_buckets_keeps_lowest() {
561        let resp: RetrieveUserQuotaResponse = serde_json::from_str(BUCKETS).unwrap();
562        let quotas = AntigravityProvider::parse_buckets(&resp);
563        let claude = quotas
564            .iter()
565            .find(|q| q.model_id == "claude-sonnet")
566            .unwrap();
567        assert_eq!(claude.remaining_fraction, Some(0.3));
568    }
569
570    #[test]
571    fn test_snapshot_orders_by_lowest() {
572        let resp: FetchAvailableModelsResponse = serde_json::from_str(MODELS).unwrap();
573        let quotas = AntigravityProvider::parse_models(&resp);
574        let snapshot = AntigravityProvider::snapshot_from_quotas(&quotas);
575        // claude 25% left → 75% used is the lowest remaining → primary.
576        assert_eq!(snapshot.primary_rate_window.unwrap().used, Some(75));
577        assert_eq!(snapshot.secondary_rate_window.unwrap().used, Some(20));
578    }
579
580    #[tokio::test]
581    async fn test_fetch_usage_with_models() {
582        let server = MockServer::start().await;
583        Mock::given(method("POST"))
584            .and(path("/v1internal:fetchAvailableModels"))
585            .respond_with(ResponseTemplate::new(200).set_body_raw(MODELS, "application/json"))
586            .mount(&server)
587            .await;
588
589        let provider = AntigravityProvider::with_base_url(&server.uri());
590        let mut ctx = ProviderContext::new();
591        ctx.config.insert("access_token".into(), "ya29-test".into());
592        let snapshot = provider.fetch_usage(&ctx).await.unwrap();
593        assert_eq!(snapshot.primary_rate_window.unwrap().used, Some(75));
594    }
595
596    #[tokio::test]
597    async fn test_fetch_usage_falls_back_to_buckets() {
598        let server = MockServer::start().await;
599        // All models full → triggers retrieveUserQuota fallback.
600        Mock::given(method("POST"))
601            .and(path("/v1internal:fetchAvailableModels"))
602            .respond_with(ResponseTemplate::new(200).set_body_raw(
603                r#"{"models":{"gemini-pro":{"displayName":"Gemini Pro","quotaInfo":{"remainingFraction":1.0}}}}"#,
604                "application/json",
605            ))
606            .mount(&server)
607            .await;
608        Mock::given(method("POST"))
609            .and(path("/v1internal:retrieveUserQuota"))
610            .respond_with(ResponseTemplate::new(200).set_body_raw(BUCKETS, "application/json"))
611            .mount(&server)
612            .await;
613
614        let provider = AntigravityProvider::with_base_url(&server.uri());
615        let mut ctx = ProviderContext::new();
616        ctx.config.insert("access_token".into(), "ya29-test".into());
617        let snapshot = provider.fetch_usage(&ctx).await.unwrap();
618        // buckets: claude 30% left → 70% used is lowest → primary.
619        assert_eq!(snapshot.primary_rate_window.unwrap().used, Some(70));
620    }
621
622    #[tokio::test]
623    async fn test_fetch_usage_401_is_auth_failed() {
624        let server = MockServer::start().await;
625        Mock::given(method("POST"))
626            .and(path("/v1internal:fetchAvailableModels"))
627            .respond_with(ResponseTemplate::new(401))
628            .mount(&server)
629            .await;
630
631        let provider = AntigravityProvider::with_base_url(&server.uri());
632        let mut ctx = ProviderContext::new();
633        ctx.config.insert("access_token".into(), "bad".into());
634        let err = provider.fetch_usage(&ctx).await.unwrap_err();
635        assert!(matches!(err, SpendPanelError::AuthFailed(_, _)));
636    }
637}