Skip to main content

lean_ctx/proxy/
usage_azure.rs

1//! Azure-specific response detection and usage normalization.
2
3use axum::http::HeaderMap;
4use serde_json::Value;
5
6/// Token usage fields reported by Azure OpenAI responses.
7#[derive(Debug, Clone, Default)]
8pub struct AzureUsage {
9    /// Tokens consumed by the prompt.
10    pub prompt_tokens: u64,
11    /// Tokens generated by the model.
12    pub completion_tokens: u64,
13    /// Prompt tokens served from cache.
14    pub cached_tokens: u64,
15    /// Total tokens reported by Azure.
16    pub total_tokens: u64,
17}
18
19/// Returns whether response headers contain an Azure-specific marker.
20pub 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
29/// Normalizes an Azure deployment or model name.
30pub 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
39/// Parses Azure token usage from a response body.
40pub 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
52/// Returns tokens attributed to Azure prompt content filtering.
53pub 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}