Skip to main content

usage_monitor_cli/provider/
zai.rs

1use async_trait::async_trait;
2use chrono::Utc;
3
4use crate::error::SpendPanelError;
5use crate::model::{PlanInfo, RateWindow, UsageSnapshot};
6use crate::provider::{ProviderContext, ProviderMetadata, UsageProvider};
7
8#[derive(Debug, serde::Deserialize)]
9struct ZaiResponse {
10    #[serde(default)]
11    code: i64,
12    #[serde(default)]
13    msg: String,
14    #[serde(default)]
15    success: bool,
16    #[serde(default)]
17    data: Option<ZaiData>,
18}
19
20#[derive(Debug, serde::Deserialize)]
21struct ZaiData {
22    #[serde(default)]
23    limits: Vec<ZaiLimit>,
24    #[serde(
25        default,
26        rename = "planName",
27        alias = "plan",
28        alias = "plan_type",
29        alias = "packageName"
30    )]
31    plan_name: Option<String>,
32}
33
34#[derive(Debug, serde::Deserialize)]
35struct ZaiLimit {
36    #[serde(rename = "type")]
37    limit_type: String,
38    #[serde(default)]
39    unit: i64,
40    #[serde(default)]
41    number: i64,
42    #[serde(default)]
43    percentage: f64,
44    /// Quota size (z.ai reports the limit under `usage`).
45    #[serde(default)]
46    usage: Option<i64>,
47    #[serde(default, rename = "currentValue")]
48    current_value: Option<i64>,
49    #[serde(default)]
50    remaining: Option<i64>,
51    #[serde(default, rename = "nextResetTime")]
52    next_reset_time: Option<i64>,
53}
54
55impl ZaiLimit {
56    /// Window length in minutes from unit code (1=days,3=hours,5=minutes,6=weeks).
57    fn window_minutes(&self) -> u32 {
58        let unit_minutes = match self.unit {
59            1 => 24 * 60,
60            3 => 60,
61            5 => 1,
62            6 => 7 * 24 * 60,
63            _ => 0,
64        };
65        (self.number.max(0) as u32).saturating_mul(unit_minutes)
66    }
67
68    /// Used percent: computed from the raw quota when present (`usage` is the
69    /// limit), else the server-provided `percentage`. Mirrors CodexBar.
70    fn used_percent(&self) -> f64 {
71        if let Some(limit) = self.usage.filter(|l| *l > 0) {
72            let used = match (self.remaining, self.current_value) {
73                (Some(remaining), Some(current)) => (limit - remaining).max(current),
74                (Some(remaining), None) => limit - remaining,
75                (None, Some(current)) => current,
76                (None, None) => return self.percentage.clamp(0.0, 100.0),
77            };
78            return ((used.max(0) as f64) / limit as f64 * 100.0).clamp(0.0, 100.0);
79        }
80        self.percentage.clamp(0.0, 100.0)
81    }
82
83    fn to_window(&self, label: &str) -> RateWindow {
84        let used = self.used_percent().round() as u64;
85        let mut w = RateWindow::new(used, 100, label.to_string(), self.window_minutes());
86        w.resets_at = self
87            .next_reset_time
88            .and_then(|ms| chrono::TimeZone::timestamp_opt(&Utc, ms / 1000, 0).single());
89        w
90    }
91}
92
93/// z.ai coding-plan quota provider (API-key auth).
94pub struct ZaiProvider {
95    metadata: ProviderMetadata,
96    base_url: Option<String>,
97}
98
99impl ZaiProvider {
100    pub fn new() -> Self {
101        Self {
102            metadata: ProviderMetadata {
103                id: "zai",
104                name: "z.ai",
105                description: "z.ai coding-plan quota monitor",
106                auth_methods: &["api_key", "env"],
107                website: Some("https://z.ai"),
108            },
109            base_url: None,
110        }
111    }
112
113    pub fn with_base_url(url: &str) -> Self {
114        let mut p = Self::new();
115        p.base_url = Some(url.to_string());
116        p
117    }
118
119    fn clean(raw: &str) -> String {
120        let mut v = raw.trim();
121        if v.len() >= 2
122            && ((v.starts_with('"') && v.ends_with('"'))
123                || (v.starts_with('\'') && v.ends_with('\'')))
124        {
125            v = &v[1..v.len() - 1];
126        }
127        v.trim().to_string()
128    }
129
130    /// Quota endpoint: explicit base_url/host config or env, else global default.
131    fn quota_url(&self, ctx: &ProviderContext) -> String {
132        if let Some(base) = self.base_url.as_deref() {
133            return format!(
134                "{}/api/monitor/usage/quota/limit",
135                base.trim_end_matches('/')
136            );
137        }
138        let host = ctx
139            .config
140            .get("base_url")
141            .or_else(|| ctx.config.get("host"))
142            .map(|s| Self::clean(s))
143            .filter(|s| !s.is_empty())
144            .or_else(|| {
145                std::env::var("Z_AI_API_HOST")
146                    .ok()
147                    .filter(|s| !s.is_empty())
148            })
149            .unwrap_or_else(|| "https://api.z.ai".to_string());
150        let host = if host.starts_with("http") {
151            host
152        } else {
153            format!("https://{}", host)
154        };
155        format!(
156            "{}/api/monitor/usage/quota/limit",
157            host.trim_end_matches('/')
158        )
159    }
160
161    fn resolve_key(ctx: &ProviderContext) -> Result<String, SpendPanelError> {
162        for key in ["api_key", "token"] {
163            if let Some(v) = ctx.config.get(key) {
164                let c = Self::clean(v);
165                if !c.is_empty() {
166                    return Ok(c);
167                }
168            }
169        }
170        if let Ok(v) = std::env::var("Z_AI_API_KEY") {
171            let c = Self::clean(&v);
172            if !c.is_empty() {
173                return Ok(c);
174            }
175        }
176        Err(SpendPanelError::AuthFailed(
177            "zai".into(),
178            "no API key in api_key/token config or Z_AI_API_KEY".into(),
179        ))
180    }
181
182    fn build_client(ctx: &ProviderContext) -> Result<reqwest::Client, SpendPanelError> {
183        reqwest::Client::builder()
184            .timeout(std::time::Duration::from_secs(ctx.timeout_secs))
185            .build()
186            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))
187    }
188
189    fn parse(body: &str) -> Result<UsageSnapshot, SpendPanelError> {
190        let resp: ZaiResponse = serde_json::from_str(body)
191            .map_err(|e| SpendPanelError::ParseError("zai".into(), e.to_string()))?;
192        if !(resp.success && resp.code == 200) {
193            return Err(SpendPanelError::ProviderError(
194                "zai".into(),
195                format!("API error (code {}): {}", resp.code, resp.msg),
196            ));
197        }
198        let data = resp.data.ok_or_else(|| {
199            SpendPanelError::ParseError("zai".into(), "missing data in response".into())
200        })?;
201
202        let mut token_limits: Vec<&ZaiLimit> = data
203            .limits
204            .iter()
205            .filter(|l| l.limit_type == "TOKENS_LIMIT")
206            .collect();
207        let time_limit = data.limits.iter().find(|l| l.limit_type == "TIME_LIMIT");
208
209        let mut snapshot = UsageSnapshot::new("zai");
210
211        // Multiple token limits: shortest window → tertiary (session), longest → primary.
212        token_limits.sort_by_key(|l| l.window_minutes());
213        if token_limits.len() >= 2 {
214            snapshot.tertiary_rate_window =
215                Some(token_limits.first().unwrap().to_window("Session tokens"));
216            snapshot.primary_rate_window = Some(token_limits.last().unwrap().to_window("Tokens"));
217        } else if let Some(only) = token_limits.first() {
218            snapshot.primary_rate_window = Some(only.to_window("Tokens"));
219        }
220
221        if let Some(time) = time_limit {
222            snapshot.secondary_rate_window = Some(time.to_window("Prompts"));
223        }
224
225        if snapshot.primary_rate_window.is_none() && snapshot.secondary_rate_window.is_none() {
226            return Err(SpendPanelError::ParseError(
227                "zai".into(),
228                "no usable limits in response".into(),
229            ));
230        }
231
232        if let Some(plan) = data.plan_name.filter(|s| !s.is_empty()) {
233            snapshot.plan = Some(PlanInfo {
234                name: plan,
235                tier: None,
236                features: Vec::new(),
237                price: None,
238                currency: None,
239                billing_period: None,
240            });
241        }
242        Ok(snapshot)
243    }
244}
245
246impl Default for ZaiProvider {
247    fn default() -> Self {
248        Self::new()
249    }
250}
251
252#[async_trait]
253impl UsageProvider for ZaiProvider {
254    fn metadata(&self) -> &ProviderMetadata {
255        &self.metadata
256    }
257
258    fn detect_credentials(&self) -> bool {
259        std::env::var("Z_AI_API_KEY")
260            .map(|v| !v.trim().is_empty())
261            .unwrap_or(false)
262    }
263
264    async fn fetch_usage(&self, ctx: &ProviderContext) -> Result<UsageSnapshot, SpendPanelError> {
265        let key = Self::resolve_key(ctx)?;
266        let client = Self::build_client(ctx)?;
267        let resp = client
268            .get(self.quota_url(ctx))
269            .header("authorization", format!("Bearer {}", key))
270            .header("accept", "application/json")
271            .send()
272            .await
273            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
274        let status = resp.status();
275        let body = resp
276            .text()
277            .await
278            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
279        if status == reqwest::StatusCode::UNAUTHORIZED || status == reqwest::StatusCode::FORBIDDEN {
280            return Err(SpendPanelError::AuthFailed(
281                "zai".into(),
282                format!("invalid API key (HTTP {})", status.as_u16()),
283            ));
284        }
285        if !status.is_success() {
286            return Err(SpendPanelError::ProviderError(
287                "zai".into(),
288                format!("HTTP {}: {}", status, body),
289            ));
290        }
291        if body.trim().is_empty() {
292            return Err(SpendPanelError::ParseError(
293                "zai".into(),
294                "empty response (check region: Global vs BigModel CN)".into(),
295            ));
296        }
297        Self::parse(&body)
298    }
299}
300
301#[cfg(test)]
302mod tests {
303    use super::*;
304    use pretty_assertions::assert_eq;
305    use wiremock::matchers::{header, method, path};
306    use wiremock::{Mock, MockServer, ResponseTemplate};
307
308    const SAMPLE: &str = r#"{
309      "code": 200, "msg": "ok", "success": true,
310      "data": {
311        "planName": "Coding Pro",
312        "limits": [
313          {"type": "TOKENS_LIMIT", "unit": 6, "number": 1, "percentage": 30, "nextResetTime": 1788000000000},
314          {"type": "TOKENS_LIMIT", "unit": 3, "number": 5, "percentage": 70},
315          {"type": "TIME_LIMIT", "unit": 3, "number": 5, "percentage": 10}
316        ]
317      }
318    }"#;
319
320    #[test]
321    fn test_metadata() {
322        assert_eq!(ZaiProvider::new().metadata().id, "zai");
323    }
324
325    #[test]
326    fn test_parse_token_window_split() {
327        let snap = ZaiProvider::parse(SAMPLE).unwrap();
328        // shortest window (5h) → tertiary, longest (1 week) → primary
329        assert_eq!(snap.primary_rate_window.unwrap().used, Some(30));
330        assert_eq!(snap.tertiary_rate_window.unwrap().used, Some(70));
331        assert_eq!(snap.secondary_rate_window.unwrap().used, Some(10));
332        assert_eq!(snap.plan.unwrap().name, "Coding Pro");
333    }
334
335    #[test]
336    fn test_used_percent_computed_from_raw_quota() {
337        // usage = limit (200), remaining 50 → used 150 → 75%.
338        let body = r#"{
339          "code": 200, "msg": "ok", "success": true,
340          "data": {"limits": [
341            {"type": "TOKENS_LIMIT", "unit": 6, "number": 1, "percentage": 10,
342             "usage": 200, "remaining": 50}
343          ]}
344        }"#;
345        let snap = ZaiProvider::parse(body).unwrap();
346        // computed (75) wins over the server `percentage` (10).
347        assert_eq!(snap.primary_rate_window.unwrap().used, Some(75));
348    }
349
350    #[test]
351    fn test_used_percent_falls_back_to_percentage() {
352        // No raw quota fields → use the server-provided percentage.
353        let body = r#"{
354          "code": 200, "msg": "ok", "success": true,
355          "data": {"limits": [
356            {"type": "TOKENS_LIMIT", "unit": 6, "number": 1, "percentage": 42}
357          ]}
358        }"#;
359        let snap = ZaiProvider::parse(body).unwrap();
360        assert_eq!(snap.primary_rate_window.unwrap().used, Some(42));
361    }
362
363    #[test]
364    fn test_api_error() {
365        let body = r#"{"code": 401, "msg": "bad token", "success": false}"#;
366        assert!(matches!(
367            ZaiProvider::parse(body).unwrap_err(),
368            SpendPanelError::ProviderError(_, _)
369        ));
370    }
371
372    #[tokio::test]
373    async fn test_fetch_success() {
374        let server = MockServer::start().await;
375        Mock::given(method("GET"))
376            .and(path("/api/monitor/usage/quota/limit"))
377            .and(header("authorization", "Bearer z"))
378            .respond_with(ResponseTemplate::new(200).set_body_raw(SAMPLE, "application/json"))
379            .mount(&server)
380            .await;
381        let provider = ZaiProvider::with_base_url(&server.uri());
382        let mut ctx = ProviderContext::new();
383        ctx.config.insert("api_key".into(), "z".into());
384        let snap = provider.fetch_usage(&ctx).await.unwrap();
385        assert_eq!(snap.primary_rate_window.unwrap().used, Some(30));
386    }
387
388    #[tokio::test]
389    async fn test_fetch_401() {
390        let server = MockServer::start().await;
391        Mock::given(method("GET"))
392            .and(path("/api/monitor/usage/quota/limit"))
393            .respond_with(ResponseTemplate::new(401))
394            .mount(&server)
395            .await;
396        let provider = ZaiProvider::with_base_url(&server.uri());
397        let mut ctx = ProviderContext::new();
398        ctx.config.insert("api_key".into(), "bad".into());
399        assert!(matches!(
400            provider.fetch_usage(&ctx).await.unwrap_err(),
401            SpendPanelError::AuthFailed(_, _)
402        ));
403    }
404}