1use 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#[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
95pub 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 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(¶ms)
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 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 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("as);
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("as);
575 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 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 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}