pg-api 0.3.13

A high-performance PostgreSQL REST API driver with rate limiting, connection pooling, and observability
use crate::auth::{extract_api_key, lookup_account};
use crate::models::AppState;
use axum::{
    extract::{Request, State},
    http::StatusCode,
    middleware::Next,
    response::Response,
};

pub async fn connection_limit_middleware(
    State(state): State<AppState>,
    request: Request,
    next: Next,
) -> Result<Response, StatusCode> {
    // Skip connection limiting for non-query endpoints
    let path = request.uri().path();
    if !path.starts_with("/v1/query")
        && !path.starts_with("/v1/batch")
        && !path.starts_with("/v1/transaction")
    {
        return Ok(next.run(request).await);
    }

    // Extract API key
    let api_key = match extract_api_key(request.headers()) {
        Ok(key) => key,
        Err(_) => {
            // Let auth middleware handle missing API key
            return Ok(next.run(request).await);
        }
    };

    // Get account info (map is keyed by hashed keys — see auth::lookup_account)
    let accounts = state.accounts.read().await;
    if let Some(account) = lookup_account(&accounts, &api_key) {
        let account_id = account.id.clone();
        let max_connections = account.max_connections;
        drop(accounts); // Release the read lock

        // Skip connection limiting if limit is 0 (unlimited)
        if max_connections == 0 {
            return Ok(next.run(request).await);
        }

        // A semaphore makes check-and-acquire atomic and releases on drop,
        // including cancellation and panics.
        let semaphore = state
            .connection_limits
            .entry(account_id)
            .or_insert_with(|| {
                std::sync::Arc::new(tokio::sync::Semaphore::new(max_connections as usize))
            })
            .clone();
        let _permit = semaphore
            .acquire_owned()
            .await
            .map_err(|_| StatusCode::SERVICE_UNAVAILABLE)?;

        Ok(next.run(request).await)
    } else {
        // Account not found, let auth middleware handle it
        Ok(next.run(request).await)
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use axum::http::{HeaderMap, HeaderValue};

    #[test]
    fn test_extract_api_key() {
        let mut headers = HeaderMap::new();

        // Test x-api-key header
        headers.insert("x-api-key", HeaderValue::from_static("test-key"));
        assert_eq!(extract_api_key(&headers).unwrap(), "test-key");

        // Test authorization header with Bearer
        headers.clear();
        headers.insert(
            "authorization",
            HeaderValue::from_static("Bearer test-key-2"),
        );
        assert_eq!(extract_api_key(&headers).unwrap(), "test-key-2");

        // Test missing key
        headers.clear();
        assert!(extract_api_key(&headers).is_err());
    }
}