use axum::{
extract::Request,
http::{HeaderMap, StatusCode},
middleware::Next,
response::Response,
};
use sha2::{Digest, Sha256};
use std::collections::HashMap;
use uuid::Uuid;
use crate::models::{Account, AppState};
pub fn hash_api_key(raw: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(raw.as_bytes());
let digest = hasher.finalize();
digest.iter().map(|b| format!("{b:02x}")).collect()
}
fn is_sha256_hex(s: &str) -> bool {
s.len() == 64 && s.bytes().all(|b| b.is_ascii_hexdigit())
}
pub fn normalize_stored_key(stored: &str) -> (String, bool) {
if is_sha256_hex(stored) {
(stored.to_ascii_lowercase(), false)
} else {
(hash_api_key(stored), true)
}
}
fn constant_time_eq(a: &str, b: &str) -> bool {
let (ab, bb) = (a.as_bytes(), b.as_bytes());
if ab.len() != bb.len() {
return false;
}
let mut diff = 0u8;
for i in 0..ab.len() {
diff |= ab[i] ^ bb[i];
}
diff == 0
}
pub fn lookup_account(accounts: &HashMap<String, Account>, presented: &str) -> Option<Account> {
let candidate = hash_api_key(presented);
accounts.iter().find_map(|(stored, account)| {
if constant_time_eq(stored, &candidate) {
let mut account = account.clone();
account.last_used = chrono::Utc::now();
Some(account)
} else {
None
}
})
}
pub async fn touch_last_used(state: &AppState, account_id: &str) {
let accounts = state.accounts.read().await;
if let Some(stored_key) = accounts
.iter()
.find(|(_, a)| a.id == account_id)
.map(|(k, _)| k.clone())
{
drop(accounts);
let mut accounts = state.accounts.write().await;
if let Some(account) = accounts.get_mut(&stored_key) {
account.last_used = chrono::Utc::now();
}
}
}
pub async fn auth_middleware(
axum::extract::State(state): axum::extract::State<AppState>,
mut request: Request,
next: Next,
) -> Result<Response, StatusCode> {
if request.uri().path() == "/health" {
return Ok(next.run(request).await);
}
let api_key = extract_api_key(request.headers()).map_err(|_| StatusCode::UNAUTHORIZED)?;
let account = {
let accounts = state.accounts.read().await;
lookup_account(&accounts, &api_key)
};
match account {
Some(account) => {
let request_id = request
.headers()
.get("x-request-id")
.and_then(|v| v.to_str().ok())
.unwrap_or("unknown");
let client_ip = request
.headers()
.get("x-forwarded-for")
.or_else(|| request.headers().get("x-real-ip"))
.and_then(|v| v.to_str().ok())
.unwrap_or("unknown");
tracing::info!(event_type = "auth_success", request_id, client_ip, account_id = %account.id, account_name = %account.name, "authenticated API request");
touch_last_used(&state, &account.id).await;
request.extensions_mut().insert(account);
Ok(next.run(request).await)
}
None => {
tracing::warn!("auth rejected: unknown api key");
Err(StatusCode::UNAUTHORIZED)
}
}
}
pub async fn request_id_middleware(mut request: Request, next: Next) -> Response {
let request_id = Uuid::new_v4().to_string();
request.headers_mut().insert(
"x-request-id",
request_id.parse().expect("uuid is a valid header value"),
);
let mut response = next.run(request).await;
if let Ok(value) = request_id.parse() {
response.headers_mut().insert("x-request-id", value);
}
response
}
pub fn extract_api_key(headers: &HeaderMap) -> Result<String, String> {
headers
.get("x-api-key")
.or_else(|| headers.get("authorization"))
.and_then(|h| h.to_str().ok())
.map(|h| {
if let Some(rest) = h.strip_prefix("Bearer ") {
rest.trim().to_string()
} else {
h.trim().to_string()
}
})
.filter(|k| !k.is_empty())
.ok_or_else(|| "Missing API key".to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_extract_api_key_from_header() {
let mut headers = HeaderMap::new();
headers.insert("x-api-key", "sk_test_123456".parse().unwrap());
let result = extract_api_key(&headers);
assert_eq!(result.unwrap(), "sk_test_123456");
}
#[test]
fn test_extract_api_key_from_bearer() {
let mut headers = HeaderMap::new();
headers.insert("authorization", "Bearer sk_test_123456".parse().unwrap());
let result = extract_api_key(&headers);
assert_eq!(result.unwrap(), "sk_test_123456");
}
#[test]
fn test_extract_api_key_missing() {
let headers = HeaderMap::new();
let result = extract_api_key(&headers);
assert!(result.is_err());
assert_eq!(result.unwrap_err(), "Missing API key");
}
#[test]
fn test_extract_api_key_from_x_api_key_priority() {
let mut headers = HeaderMap::new();
headers.insert("x-api-key", "key_from_x_api_key".parse().unwrap());
headers.insert("authorization", "Bearer key_from_auth".parse().unwrap());
let result = extract_api_key(&headers);
assert_eq!(result.unwrap(), "key_from_x_api_key");
}
#[test]
fn test_extract_api_key_case_insensitive_header() {
let mut headers = HeaderMap::new();
headers.insert("X-API-Key", "sk_test_uppercase".parse().unwrap());
let result = extract_api_key(&headers);
assert_eq!(result.unwrap(), "sk_test_uppercase");
}
#[test]
fn test_extract_api_key_bearer_with_extra_spaces() {
let mut headers = HeaderMap::new();
headers.insert("authorization", "Bearer sk_test_spaces".parse().unwrap());
let result = extract_api_key(&headers);
assert_eq!(result.unwrap(), "sk_test_spaces");
}
#[test]
fn test_hash_is_stable_hex() {
let h1 = hash_api_key("secret");
let h2 = hash_api_key("secret");
assert_eq!(h1, h2);
assert_eq!(h1.len(), 64);
assert!(is_sha256_hex(&h1));
assert_ne!(h1, hash_api_key("other"));
}
#[test]
fn test_normalize_keeps_hashes_hashes_plaintext() {
let (kept, migrated) = normalize_stored_key(&hash_api_key("abc"));
assert!(!migrated);
assert_eq!(kept, hash_api_key("abc"));
let (hashed, migrated) = normalize_stored_key("legacy-plaintext-key");
assert!(migrated);
assert_eq!(hashed, hash_api_key("legacy-plaintext-key"));
}
#[test]
fn test_lookup_account_matches_raw_key() {
use crate::models::{AccountRole, DatabaseAccess};
use chrono::Utc;
let raw = "my-raw-key";
let mut map = HashMap::new();
map.insert(
hash_api_key(raw),
Account {
id: "acc_1".to_string(),
name: "t".to_string(),
api_key: hash_api_key(raw),
instance_id: "default".to_string(),
databases: vec![DatabaseAccess {
database: "db".to_string(),
username: "u".to_string(),
password: "p".to_string(),
permissions: vec![],
}],
role: AccountRole::Readonly,
created_at: Utc::now(),
last_used: Utc::now(),
rate_limit: 10,
max_connections: 2,
notes: None,
},
);
assert!(lookup_account(&map, raw).is_some());
assert!(lookup_account(&map, "wrong").is_none());
}
}