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> {
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);
}
let api_key = match extract_api_key(request.headers()) {
Ok(key) => key,
Err(_) => {
return Ok(next.run(request).await);
}
};
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);
if max_connections == 0 {
return Ok(next.run(request).await);
}
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 {
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();
headers.insert("x-api-key", HeaderValue::from_static("test-key"));
assert_eq!(extract_api_key(&headers).unwrap(), "test-key");
headers.clear();
headers.insert(
"authorization",
HeaderValue::from_static("Bearer test-key-2"),
);
assert_eq!(extract_api_key(&headers).unwrap(), "test-key-2");
headers.clear();
assert!(extract_api_key(&headers).is_err());
}
}