lean_ctx/proxy/
usage_azure.rs1use axum::http::HeaderMap;
4use serde_json::Value;
5
6#[derive(Debug, Clone, Default)]
8pub struct AzureUsage {
9 pub prompt_tokens: u64,
11 pub completion_tokens: u64,
13 pub cached_tokens: u64,
15 pub total_tokens: u64,
17}
18
19pub fn is_azure_response(headers: &HeaderMap) -> bool {
21 headers.contains_key("x-ms-region")
22 || headers.contains_key("x-ms-rai-invoked")
23 || headers
24 .keys()
25 .map(http::HeaderName::as_str)
26 .any(|name| name.starts_with("x-ms-") && name.contains("azure"))
27}
28
29pub fn normalize_azure_model(model: &str) -> String {
31 let model = model.trim();
32 if model.is_empty() {
33 "unknown".to_string()
34 } else {
35 model.to_string()
36 }
37}
38
39pub fn parse_azure_usage(body: &Value) -> Option<AzureUsage> {
41 let usage = body.get("usage")?;
42 Some(AzureUsage {
43 prompt_tokens: usage["prompt_tokens"].as_u64().unwrap_or_default(),
44 completion_tokens: usage["completion_tokens"].as_u64().unwrap_or_default(),
45 cached_tokens: usage["prompt_tokens_details"]["cached_tokens"]
46 .as_u64()
47 .unwrap_or_default(),
48 total_tokens: usage["total_tokens"].as_u64().unwrap_or_default(),
49 })
50}
51
52pub fn content_filter_overhead(body: &Value) -> u64 {
54 body["usage"]["prompt_filter_results"]
55 .as_u64()
56 .unwrap_or_default()
57}
58
59#[cfg(test)]
60mod tests {
61 use super::{
62 content_filter_overhead, is_azure_response, normalize_azure_model, parse_azure_usage,
63 };
64 use axum::http::HeaderMap;
65 use serde_json::json;
66
67 #[test]
68 fn is_azure_by_region() {
69 let mut headers = HeaderMap::new();
70 headers.insert("x-ms-region", "westus".parse().expect("valid header"));
71 assert!(is_azure_response(&headers));
72 }
73
74 #[test]
75 fn is_azure_by_rai() {
76 let mut headers = HeaderMap::new();
77 headers.insert("x-ms-rai-invoked", "true".parse().expect("valid header"));
78 assert!(is_azure_response(&headers));
79 }
80
81 #[test]
82 fn not_azure_without_headers() {
83 assert!(!is_azure_response(&HeaderMap::new()));
84 }
85
86 #[test]
87 fn parse_standard_usage() {
88 let body = json!({"usage": {
89 "prompt_tokens": 100, "completion_tokens": 50, "total_tokens": 150
90 }});
91 let usage = parse_azure_usage(&body).expect("usage present");
92 assert_eq!(usage.prompt_tokens, 100);
93 assert_eq!(usage.completion_tokens, 50);
94 assert_eq!(usage.total_tokens, 150);
95 }
96
97 #[test]
98 fn parse_with_cache() {
99 let body = json!({"usage": {
100 "prompt_tokens_details": {"cached_tokens": 40}
101 }});
102 assert_eq!(
103 parse_azure_usage(&body)
104 .expect("usage present")
105 .cached_tokens,
106 40
107 );
108 }
109
110 #[test]
111 fn missing_usage_returns_none() {
112 assert!(parse_azure_usage(&json!({})).is_none());
113 }
114
115 #[test]
116 fn normalize_model_passthrough() {
117 assert_eq!(normalize_azure_model("gpt-4o"), "gpt-4o");
118 assert_eq!(normalize_azure_model(" gpt-4o "), "gpt-4o");
119 assert_eq!(normalize_azure_model(""), "unknown");
120 }
121
122 #[test]
123 fn reads_content_filter_overhead() {
124 let body = json!({"usage": {"prompt_filter_results": 7}});
125 assert_eq!(content_filter_overhead(&body), 7);
126 assert_eq!(content_filter_overhead(&json!({})), 0);
127 }
128}