Skip to main content

usage_monitor_cli/provider/
groq.rs

1use async_trait::async_trait;
2
3use crate::error::SpendPanelError;
4use crate::model::{RateWindow, RateWindowStatus, UsageSnapshot};
5use crate::provider::{ProviderContext, ProviderMetadata, UsageProvider};
6
7#[derive(Debug, serde::Deserialize)]
8struct PrometheusResponse {
9    status: String,
10    data: Option<PrometheusPayload>,
11    error: Option<String>,
12}
13
14#[derive(Debug, serde::Deserialize)]
15struct PrometheusPayload {
16    result: Vec<PrometheusSeries>,
17}
18
19#[derive(Debug, serde::Deserialize)]
20struct PrometheusSeries {
21    value: Option<Vec<PrometheusValue>>,
22}
23
24#[derive(Debug, serde::Deserialize)]
25#[serde(untagged)]
26enum PrometheusValue {
27    Number(f64),
28    String(String),
29}
30
31impl PrometheusValue {
32    fn as_f64(&self) -> Option<f64> {
33        match self {
34            Self::Number(n) => Some(*n),
35            Self::String(s) => s.parse().ok(),
36        }
37    }
38}
39
40#[derive(Debug, Clone, PartialEq)]
41struct GroqUsage {
42    request_rate_per_second: f64,
43    input_token_rate_per_second: f64,
44    output_token_rate_per_second: f64,
45    prompt_cache_hit_rate_per_second: f64,
46}
47
48impl GroqUsage {
49    fn requests_per_minute(&self) -> f64 {
50        self.request_rate_per_second * 60.0
51    }
52
53    fn tokens_per_minute(&self) -> f64 {
54        (self.input_token_rate_per_second + self.output_token_rate_per_second) * 60.0
55    }
56
57    fn cache_hits_per_minute(&self) -> f64 {
58        self.prompt_cache_hit_rate_per_second * 60.0
59    }
60}
61
62/// GroqCloud Prometheus metrics provider.
63pub struct GroqProvider {
64    metadata: ProviderMetadata,
65    /// Base URL override for tests.
66    base_url: Option<String>,
67}
68
69impl GroqProvider {
70    pub fn new() -> Self {
71        Self {
72            metadata: ProviderMetadata {
73                id: "groq",
74                name: "GroqCloud",
75                description: "GroqCloud Prometheus metrics monitor",
76                auth_methods: &["api_key", "env"],
77                website: Some("https://console.groq.com"),
78            },
79            base_url: None,
80        }
81    }
82
83    pub fn with_base_url(url: &str) -> Self {
84        let mut p = Self::new();
85        p.base_url = Some(url.to_string());
86        p
87    }
88
89    fn clean(raw: &str) -> String {
90        let mut value = raw.trim();
91        if value.len() >= 2
92            && ((value.starts_with('"') && value.ends_with('"'))
93                || (value.starts_with('\'') && value.ends_with('\'')))
94        {
95            value = &value[1..value.len() - 1];
96        }
97        value.trim().to_string()
98    }
99
100    fn resolve_api_key(ctx: &ProviderContext) -> Result<String, SpendPanelError> {
101        for key in ["api_key", "token"] {
102            if let Some(value) = ctx.config.get(key) {
103                let cleaned = Self::clean(value);
104                if !cleaned.is_empty() {
105                    return Ok(cleaned);
106                }
107            }
108        }
109        for env in ["GROQ_API_KEY", "GROQ_TOKEN"] {
110            if let Ok(value) = std::env::var(env) {
111                let cleaned = Self::clean(&value);
112                if !cleaned.is_empty() {
113                    return Ok(cleaned);
114                }
115            }
116        }
117        Err(SpendPanelError::AuthFailed(
118            "groq".into(),
119            "no API key found in config, token, GROQ_API_KEY, or GROQ_TOKEN".into(),
120        ))
121    }
122
123    fn api_base(&self, ctx: &ProviderContext) -> String {
124        let configured = ctx
125            .config
126            .get("api_url")
127            .or_else(|| ctx.config.get("base_url"))
128            .map(String::as_str)
129            .filter(|v| !v.is_empty())
130            .map(Self::clean)
131            .or_else(|| std::env::var("GROQ_API_URL").ok().map(|v| Self::clean(&v)))
132            .or_else(|| self.base_url.clone())
133            .unwrap_or_else(|| "https://api.groq.com/v1".into());
134
135        let base = if configured.starts_with("http://") || configured.starts_with("https://") {
136            configured
137        } else {
138            format!("https://{}", configured)
139        };
140        format!("{}/metrics/prometheus", base.trim_end_matches('/'))
141    }
142
143    fn build_client(ctx: &ProviderContext) -> Result<reqwest::Client, SpendPanelError> {
144        reqwest::Client::builder()
145            .timeout(std::time::Duration::from_secs(ctx.timeout_secs))
146            .build()
147            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))
148    }
149
150    fn parse_scalar(body: &str) -> Result<f64, SpendPanelError> {
151        let decoded: PrometheusResponse = serde_json::from_str(body)
152            .map_err(|e| SpendPanelError::ParseError("groq".into(), e.to_string()))?;
153        if decoded.status != "success" {
154            return Err(SpendPanelError::ProviderError(
155                "groq".into(),
156                decoded.error.unwrap_or_else(|| "query failed".into()),
157            ));
158        }
159        Ok(decoded
160            .data
161            .map(|data| {
162                data.result
163                    .iter()
164                    .filter_map(|series| series.value.as_ref())
165                    .filter_map(|values| values.last())
166                    .filter_map(PrometheusValue::as_f64)
167                    .sum()
168            })
169            .unwrap_or(0.0))
170    }
171
172    async fn query_scalar(
173        client: &reqwest::Client,
174        base_url: &str,
175        api_key: &str,
176        query: &str,
177    ) -> Result<f64, SpendPanelError> {
178        let resp = client
179            .get(format!("{}/api/v1/query", base_url))
180            .query(&[("query", query)])
181            .header("Authorization", format!("Bearer {}", api_key))
182            .header("Accept", "application/json")
183            .send()
184            .await
185            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
186
187        let status = resp.status();
188        let body = resp
189            .text()
190            .await
191            .map_err(|e| SpendPanelError::NetworkError(e.to_string()))?;
192        if status == reqwest::StatusCode::UNAUTHORIZED || status == reqwest::StatusCode::FORBIDDEN {
193            return Err(SpendPanelError::AuthFailed(
194                "groq".into(),
195                format!("metrics access denied (HTTP {})", status.as_u16()),
196            ));
197        }
198        if !status.is_success() {
199            return Err(SpendPanelError::ProviderError(
200                "groq".into(),
201                format!("HTTP {}: {}", status, body),
202            ));
203        }
204        Self::parse_scalar(&body)
205    }
206
207    async fn fetch_usage_data(
208        &self,
209        client: &reqwest::Client,
210        api_key: &str,
211        ctx: &ProviderContext,
212    ) -> Result<GroqUsage, SpendPanelError> {
213        let base_url = self.api_base(ctx);
214        let (requests, input_tokens, output_tokens, cache_hits) = tokio::try_join!(
215            Self::query_scalar(
216                client,
217                &base_url,
218                api_key,
219                "sum(model_project_id_status_code:requests:rate5m)"
220            ),
221            Self::query_scalar(
222                client,
223                &base_url,
224                api_key,
225                "sum(model_project_id:tokens_in:rate5m)"
226            ),
227            Self::query_scalar(
228                client,
229                &base_url,
230                api_key,
231                "sum(model_project_id:tokens_out:rate5m)"
232            ),
233            Self::query_scalar(
234                client,
235                &base_url,
236                api_key,
237                "sum(model_project_id:prompt_cache_hits:rate5m)"
238            ),
239        )?;
240        Ok(GroqUsage {
241            request_rate_per_second: requests,
242            input_token_rate_per_second: input_tokens,
243            output_token_rate_per_second: output_tokens,
244            prompt_cache_hit_rate_per_second: cache_hits,
245        })
246    }
247
248    fn format_decimal(value: f64) -> String {
249        if value >= 100.0 {
250            format!("{:.0}", value)
251        } else if value >= 10.0 {
252            format!("{:.1}", value)
253        } else {
254            format!("{:.2}", value)
255        }
256    }
257
258    fn zero_window(label: impl Into<String>) -> RateWindow {
259        RateWindow {
260            label: label.into(),
261            window_minutes: 5,
262            usage_ratio: 0.0,
263            limit: None,
264            used: None,
265            remaining: None,
266            resets_at: None,
267            status: RateWindowStatus::Normal,
268        }
269    }
270
271    fn snapshot_from_usage(usage: GroqUsage) -> UsageSnapshot {
272        let mut snapshot = UsageSnapshot::new("groq");
273        snapshot.primary_rate_window = Some(Self::zero_window(format!(
274            "Requests {} req/min",
275            Self::format_decimal(usage.requests_per_minute())
276        )));
277        snapshot.secondary_rate_window = Some(Self::zero_window(format!(
278            "Tokens {} tok/min",
279            Self::format_decimal(usage.tokens_per_minute())
280        )));
281        if usage.prompt_cache_hit_rate_per_second > 0.0 {
282            snapshot.tertiary_rate_window = Some(Self::zero_window(format!(
283                "Cache {} cache/min",
284                Self::format_decimal(usage.cache_hits_per_minute())
285            )));
286        }
287        snapshot
288    }
289}
290
291impl Default for GroqProvider {
292    fn default() -> Self {
293        Self::new()
294    }
295}
296
297#[async_trait]
298impl UsageProvider for GroqProvider {
299    fn metadata(&self) -> &ProviderMetadata {
300        &self.metadata
301    }
302
303    fn detect_credentials(&self) -> bool {
304        ["GROQ_API_KEY", "GROQ_TOKEN"]
305            .iter()
306            .any(|env| std::env::var(env).is_ok_and(|v| !Self::clean(&v).is_empty()))
307    }
308
309    async fn fetch_usage(&self, ctx: &ProviderContext) -> Result<UsageSnapshot, SpendPanelError> {
310        let api_key = Self::resolve_api_key(ctx)?;
311        let client = Self::build_client(ctx)?;
312        let usage = self.fetch_usage_data(&client, &api_key, ctx).await?;
313        Ok(Self::snapshot_from_usage(usage))
314    }
315}
316
317#[cfg(test)]
318mod tests {
319    use super::*;
320    use pretty_assertions::assert_eq;
321    use wiremock::matchers::{header, method, path, query_param};
322    use wiremock::{Mock, MockServer, ResponseTemplate};
323
324    const SUCCESS: &str = r#"{
325      "status":"success",
326      "data":{"result":[{"value":[1710000000,"2.5"]},{"value":[1710000000,"1.5"]}]}
327    }"#;
328
329    #[test]
330    fn test_provider_metadata() {
331        let provider = GroqProvider::new();
332        let meta = provider.metadata();
333        assert_eq!(meta.id, "groq");
334        assert_eq!(meta.name, "GroqCloud");
335        assert!(meta.auth_methods.contains(&"api_key"));
336    }
337
338    #[test]
339    fn test_parse_prometheus_scalar_response() {
340        assert_eq!(GroqProvider::parse_scalar(SUCCESS).unwrap(), 4.0);
341    }
342
343    #[test]
344    fn test_parse_prometheus_error_response() {
345        let err = GroqProvider::parse_scalar(r#"{"status":"error","error":"nope"}"#).unwrap_err();
346        assert!(matches!(err, SpendPanelError::ProviderError(_, _)));
347    }
348
349    #[test]
350    fn test_snapshot_maps_rates_to_windows() {
351        let snapshot = GroqProvider::snapshot_from_usage(GroqUsage {
352            request_rate_per_second: 2.0,
353            input_token_rate_per_second: 100.0,
354            output_token_rate_per_second: 50.0,
355            prompt_cache_hit_rate_per_second: 3.0,
356        });
357        assert_eq!(
358            snapshot.primary_rate_window.unwrap().label,
359            "Requests 120 req/min"
360        );
361        assert_eq!(
362            snapshot.secondary_rate_window.unwrap().label,
363            "Tokens 9000 tok/min"
364        );
365        assert_eq!(
366            snapshot.tertiary_rate_window.unwrap().label,
367            "Cache 180 cache/min"
368        );
369    }
370
371    #[tokio::test]
372    async fn test_fetch_usage_success() {
373        let server = MockServer::start().await;
374        for query in [
375            "sum(model_project_id_status_code:requests:rate5m)",
376            "sum(model_project_id:tokens_in:rate5m)",
377            "sum(model_project_id:tokens_out:rate5m)",
378            "sum(model_project_id:prompt_cache_hits:rate5m)",
379        ] {
380            Mock::given(method("GET"))
381                .and(path("/v1/metrics/prometheus/api/v1/query"))
382                .and(query_param("query", query))
383                .and(header("authorization", "Bearer gsk-test"))
384                .and(header("accept", "application/json"))
385                .respond_with(ResponseTemplate::new(200).set_body_raw(SUCCESS, "application/json"))
386                .mount(&server)
387                .await;
388        }
389
390        let provider = GroqProvider::with_base_url(&format!("{}/v1", server.uri()));
391        let snapshot = provider
392            .fetch_usage(&ProviderContext::with_api_key("gsk-test"))
393            .await
394            .unwrap();
395        assert_eq!(
396            snapshot.primary_rate_window.unwrap().label,
397            "Requests 240 req/min"
398        );
399        assert_eq!(
400            snapshot.secondary_rate_window.unwrap().label,
401            "Tokens 480 tok/min"
402        );
403        assert_eq!(
404            snapshot.tertiary_rate_window.unwrap().label,
405            "Cache 240 cache/min"
406        );
407    }
408
409    #[tokio::test]
410    async fn test_fetch_usage_401_is_auth_failed() {
411        let server = MockServer::start().await;
412        Mock::given(method("GET"))
413            .and(path("/v1/metrics/prometheus/api/v1/query"))
414            .respond_with(ResponseTemplate::new(401))
415            .mount(&server)
416            .await;
417
418        let provider = GroqProvider::with_base_url(&format!("{}/v1", server.uri()));
419        let err = provider
420            .fetch_usage(&ProviderContext::with_api_key("bad"))
421            .await
422            .unwrap_err();
423        assert!(matches!(err, SpendPanelError::AuthFailed(_, _)));
424    }
425}