usage_monitor_cli/provider/
groq.rs1use 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
62pub struct GroqProvider {
64 metadata: ProviderMetadata,
65 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}