pg-api 0.3.7

A high-performance PostgreSQL REST API driver with rate limiting, connection pooling, and observability
//! Unified authentication module (v0.3).
//!
//! Single source of truth for API key handling:
//! - extraction from headers (used by auth, rate-limit and connection-limit layers)
//! - SHA-256 hashing of keys at rest (migration-compatible with legacy plaintext)
//! - request authentication middleware with no secret material in logs
//!
//! [`SecretProvider`] is intentionally minimal in v0.3 (`EnvOrFile` semantics):
//! keys live in `accounts.json` as hashes. A future Vault/KMS/SOPS provider
//! only needs to supply the same normalized hashes.

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};

/// Compute the SHA-256 hex digest of a presented (raw) API key.
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()
}

/// Returns true if `s` already looks like a SHA-256 hex digest.
fn is_sha256_hex(s: &str) -> bool {
    s.len() == 64 && s.bytes().all(|b| b.is_ascii_hexdigit())
}

/// Normalize a stored key for in-memory comparison.
///
/// - Already-hashed values (64 hex chars) are kept as-is (lowercased).
/// - Legacy plaintext values are hashed on load so that comparison is
///   always hash-vs-hash. Operators should re-save/rotate these entries:
///   a warning is emitted by the caller on this path.
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)
    }
}

/// Constant-time equality for fixed-length digests (avoids timing oracles).
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
}

/// Look up the account matching a presented raw API key.
/// The map must be keyed by normalized (hashed) keys — see `load_accounts`.
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
        }
    })
}

/// Persist the `last_used` timestamp back into the shared map.
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> {
    // Skip auth for health check only. Every other route (including /docs
    // and /openapi.json) requires a valid key: docs embed a live console.
    if request.uri().path() == "/health" {
        return Ok(next.run(request).await);
    }

    let api_key = extract_api_key(request.headers()).map_err(|_| StatusCode::UNAUTHORIZED)?;

    // Read-lock only: last_used is updated via a short follow-up write.
    let account = {
        let accounts = state.accounts.read().await;
        lookup_account(&accounts, &api_key)
    };

    match account {
        Some(account) => {
            touch_last_used(&state, &account.id).await;
            request.extensions_mut().insert(account);
            Ok(next.run(request).await)
        }
        None => {
            // Deliberately no key material, no account counts in logs.
            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();
    // Propagate to the downstream handler AND back to the caller.
    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
}

/// Extract the raw API key from `x-api-key` (preferred) or
/// `Authorization: Bearer <key>`. Surrounding whitespace is trimmed so that
/// `Bearer  <key>` does not produce a phantom mismatch.
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() {
        // Quando ambos headers estão presentes, x-api-key tem prioridade
        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());

        // v0.3: whitespace é normalizado (antes preservava o espaço e falhava o match)
        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());
    }
}