use axum::http::HeaderMap;
use serde_json::Value;
#[derive(Debug, Clone, Default)]
pub struct AzureUsage {
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub cached_tokens: u64,
pub total_tokens: u64,
}
pub fn is_azure_response(headers: &HeaderMap) -> bool {
headers.contains_key("x-ms-region")
|| headers.contains_key("x-ms-rai-invoked")
|| headers
.keys()
.map(http::HeaderName::as_str)
.any(|name| name.starts_with("x-ms-") && name.contains("azure"))
}
pub fn normalize_azure_model(model: &str) -> String {
let model = model.trim();
if model.is_empty() {
"unknown".to_string()
} else {
model.to_string()
}
}
pub fn parse_azure_usage(body: &Value) -> Option<AzureUsage> {
let usage = body.get("usage")?;
Some(AzureUsage {
prompt_tokens: usage["prompt_tokens"].as_u64().unwrap_or_default(),
completion_tokens: usage["completion_tokens"].as_u64().unwrap_or_default(),
cached_tokens: usage["prompt_tokens_details"]["cached_tokens"]
.as_u64()
.unwrap_or_default(),
total_tokens: usage["total_tokens"].as_u64().unwrap_or_default(),
})
}
pub fn content_filter_overhead(body: &Value) -> u64 {
body["usage"]["prompt_filter_results"]
.as_u64()
.unwrap_or_default()
}
#[cfg(test)]
mod tests {
use super::{
content_filter_overhead, is_azure_response, normalize_azure_model, parse_azure_usage,
};
use axum::http::HeaderMap;
use serde_json::json;
#[test]
fn is_azure_by_region() {
let mut headers = HeaderMap::new();
headers.insert("x-ms-region", "westus".parse().expect("valid header"));
assert!(is_azure_response(&headers));
}
#[test]
fn is_azure_by_rai() {
let mut headers = HeaderMap::new();
headers.insert("x-ms-rai-invoked", "true".parse().expect("valid header"));
assert!(is_azure_response(&headers));
}
#[test]
fn not_azure_without_headers() {
assert!(!is_azure_response(&HeaderMap::new()));
}
#[test]
fn parse_standard_usage() {
let body = json!({"usage": {
"prompt_tokens": 100, "completion_tokens": 50, "total_tokens": 150
}});
let usage = parse_azure_usage(&body).expect("usage present");
assert_eq!(usage.prompt_tokens, 100);
assert_eq!(usage.completion_tokens, 50);
assert_eq!(usage.total_tokens, 150);
}
#[test]
fn parse_with_cache() {
let body = json!({"usage": {
"prompt_tokens_details": {"cached_tokens": 40}
}});
assert_eq!(
parse_azure_usage(&body)
.expect("usage present")
.cached_tokens,
40
);
}
#[test]
fn missing_usage_returns_none() {
assert!(parse_azure_usage(&json!({})).is_none());
}
#[test]
fn normalize_model_passthrough() {
assert_eq!(normalize_azure_model("gpt-4o"), "gpt-4o");
assert_eq!(normalize_azure_model(" gpt-4o "), "gpt-4o");
assert_eq!(normalize_azure_model(""), "unknown");
}
#[test]
fn reads_content_filter_overhead() {
let body = json!({"usage": {"prompt_filter_results": 7}});
assert_eq!(content_filter_overhead(&body), 7);
assert_eq!(content_filter_overhead(&json!({})), 0);
}
}