use crate::VecboostState;
use crate::audit::AuditLogger;
use crate::auth::{GarrisonUtil, User};
use crate::config::app::AuthConfig;
use axum::{
extract::{ConnectInfo, Request, State},
http::{HeaderMap, StatusCode},
middleware::Next,
response::Response,
};
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
#[derive(Clone)]
pub struct AuthContext {
pub user: User,
pub token: String,
}
const PUBLIC_PATHS: &[&str] = &["/health", "/api/1/auth/login", "/api/1/auth/refresh"];
const ADMIN_PATH_PREFIXES: &[&str] = &[
"/api/1/model/",
"/api/1/embed/file",
"/model/",
"/embed/file",
];
fn requires_admin(path: &str) -> bool {
ADMIN_PATH_PREFIXES.iter().any(|p| path.starts_with(p))
}
async fn current_token_is_admin() -> bool {
matches!(GarrisonUtil::has_role("admin").await, Ok(true))
}
#[cfg_attr(not(feature = "grpc"), allow(dead_code))]
pub async fn current_token_is_admin_pub() -> bool {
current_token_is_admin().await
}
fn extract_client_ip(
headers: &HeaderMap,
connect_info: Option<SocketAddr>,
trusted_proxies: &[String],
) -> Option<IpAddr> {
let peer_ip = connect_info.map(|sa| sa.ip());
let xff_trusted = if trusted_proxies.is_empty() {
false
} else {
match peer_ip {
Some(ip) => crate::rate_limit::is_ip_whitelisted(&ip.to_string(), trusted_proxies),
None => false,
}
};
if xff_trusted
&& let Some(xff_ip) = headers
.get("x-forwarded-for")
.or_else(|| headers.get("x-real-ip"))
.and_then(|h| h.to_str().ok())
.and_then(|s| s.split(',').next())
.and_then(|s| s.trim().parse().ok())
{
return Some(xff_ip);
}
peer_ip
}
pub async fn auth_middleware(
State(audit_logger): State<Option<Arc<AuditLogger>>>,
State(auth_config): State<AuthConfig>,
request: Request,
next: Next,
) -> Result<Response, StatusCode> {
static EMPTY_PROXIES_WARN: std::sync::Once = std::sync::Once::new();
if auth_config.trusted_proxies.is_empty() {
EMPTY_PROXIES_WARN.call_once(|| {
log::info!(
"trusted_proxies is empty — X-Forwarded-For is ignored; the direct peer IP is \
used for rate limiting and audit. Behind a reverse proxy, configure \
trusted_proxies (e.g. [\"10.0.0.0/8\"]) to honor forwarded headers."
);
});
}
let path = request.uri().path();
if PUBLIC_PATHS.contains(&path) {
return Ok(next.run(request).await);
}
let auth_header = request
.headers()
.get("authorization")
.and_then(|h| h.to_str().ok());
let connect_info = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ci| ci.0);
let ip = extract_client_ip(
request.headers(),
connect_info,
&auth_config.trusted_proxies,
);
let token = match auth_header {
Some(header) if header.starts_with("Bearer ") => header[7..].to_string(),
_ => {
if let Some(ref logger) = audit_logger {
logger.log_unauthorized_access(ip.map(|i| i.to_string()), path);
}
return Ok(unauthorized_response("auth-credentials-missing"));
}
};
match GarrisonUtil::get_login_id_by_token(&token).await {
Ok(Some(login_id)) => {
let user = User {
username: login_id.clone(),
role: String::new(), permissions: vec![],
};
if requires_admin(path) {
let is_admin =
garrison::stp::with_current_token(token.clone(), current_token_is_admin())
.await;
if !is_admin {
if let Some(ref logger) = audit_logger {
logger.log_unauthorized_access(ip.map(|i| i.to_string()), path);
}
return Ok(forbidden_response());
}
}
let mut request = request;
request.extensions_mut().insert(AuthContext {
user,
token: token.clone(),
});
Ok(garrison::stp::with_current_token(token, next.run(request)).await)
}
_ => {
if let Some(ref logger) = audit_logger {
logger.log_unauthorized_access(ip.map(|i| i.to_string()), path);
}
Ok(unauthorized_response("auth-invalid-token"))
}
}
}
fn forbidden_response() -> Response {
use axum::http::header;
let body = serde_json::json!({
"success": false,
"error": {
"code": "FORBIDDEN",
"message": crate::i18n::tr("auth-admin-required"),
}
});
let mut resp = Response::new(axum::body::Body::from(body.to_string()));
*resp.status_mut() = StatusCode::FORBIDDEN;
resp.headers_mut().insert(
header::CONTENT_TYPE,
header::HeaderValue::from_static("application/json"),
);
resp
}
fn unauthorized_response(message_key: &str) -> Response {
use axum::http::header;
let body = serde_json::json!({
"type": "AuthenticationRequired",
"message": crate::i18n::tr(message_key),
"field": null,
"value": null,
});
let mut resp = Response::new(axum::body::Body::from(body.to_string()));
*resp.status_mut() = StatusCode::UNAUTHORIZED;
resp.headers_mut().insert(
header::CONTENT_TYPE,
header::HeaderValue::from_static("application/json"),
);
resp
}
fn rate_limited_response() -> Response {
use axum::http::header;
let body = serde_json::json!({
"type": "RateLimitExceeded",
"message": crate::i18n::tr("rate-limit-exceeded"),
"field": null,
"value": null,
});
let mut resp = Response::new(axum::body::Body::from(body.to_string()));
*resp.status_mut() = StatusCode::TOO_MANY_REQUESTS;
resp.headers_mut().insert(
header::CONTENT_TYPE,
header::HeaderValue::from_static("application/json"),
);
resp
}
pub async fn optional_auth_middleware(
headers: HeaderMap,
mut request: Request,
next: Next,
) -> Response {
if let Some(auth_header) = headers.get("authorization")
&& let Ok(auth_str) = auth_header.to_str()
&& let Some(token) = auth_str.strip_prefix("Bearer ")
&& let Ok(Some(login_id)) = GarrisonUtil::get_login_id_by_token(token).await
{
let user = User {
username: login_id,
role: String::new(),
permissions: vec![],
};
request.extensions_mut().insert(AuthContext {
user,
token: token.to_string(),
});
return garrison::stp::with_current_token(token.to_string(), next.run(request)).await;
}
next.run(request).await
}
pub async fn require_permission_middleware(
permission: &'static str,
request: Request,
next: Next,
) -> Result<Response, StatusCode> {
let _auth_context = request
.extensions()
.get::<AuthContext>()
.ok_or(StatusCode::UNAUTHORIZED)?;
match GarrisonUtil::has_permission(permission).await {
Ok(true) => Ok(next.run(request).await),
_ => Err(StatusCode::FORBIDDEN),
}
}
pub async fn require_role_middleware(request: Request, next: Next) -> Result<Response, StatusCode> {
let _auth_context = request
.extensions()
.get::<AuthContext>()
.ok_or(StatusCode::UNAUTHORIZED)?;
match GarrisonUtil::has_role("admin").await {
Ok(true) => Ok(next.run(request).await),
_ => Err(StatusCode::FORBIDDEN),
}
}
pub async fn auth_rate_limit_middleware(
State(state): State<VecboostState>,
State(auth_config): State<AuthConfig>,
request: Request,
next: Next,
) -> Result<Response, StatusCode> {
let rate_limit_enabled = state
.kit
.config::<crate::registry::RateLimitEnabled>()
.map(|c| c.0)
.unwrap_or_else(|_| {
log::warn!(
"RateLimitEnabled not registered, rate limiting is disabled. \
This may leave endpoints unprotected in production."
);
false
});
if !rate_limit_enabled {
return Ok(next.run(request).await);
}
let path = request.uri().path();
if path == "/health" || path == "/metrics" {
return Ok(next.run(request).await);
}
let ip_whitelist = state
.kit
.require::<crate::registry::IpWhitelistModule>()
.map_err(|e| {
log::error!("IpWhitelistModule not registered: {}", e);
StatusCode::INTERNAL_SERVER_ERROR
})?;
let connect_info = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ci| ci.0);
let ip = extract_client_ip(
request.headers(),
connect_info,
&auth_config.trusted_proxies,
)
.map(|i| i.to_string())
.unwrap_or_else(|| "unknown".to_string());
if crate::rate_limit::is_ip_whitelisted(&ip, &ip_whitelist) {
return Ok(next.run(request).await);
}
let rate_limiter = state
.kit
.require::<crate::registry::RateLimitModule>()
.map_err(|e| {
log::error!("RateLimitModule not registered: {}", e);
StatusCode::INTERNAL_SERVER_ERROR
})?;
let context = crate::rate_limit::RequestContext {
client_ip: Some(ip.clone()),
path: request.uri().path().to_string(),
method: request.method().to_string(),
..Default::default()
};
let decision = rate_limiter.check_rate_limit_detailed(&context).await;
let allowed = decision.allowed;
let headers_enabled = state
.kit
.config::<crate::registry::RateLimitHeadersEnabled>()
.map(|c| c.0)
.unwrap_or(false);
if let Ok(prom_collector) = state
.kit
.require::<crate::registry::PrometheusCollectorModule>()
&& let Some(prom) = prom_collector.as_ref()
{
if allowed {
prom.record_rate_limit_allowed("ip");
} else {
prom.record_rate_limit_denied("ip");
}
}
if !allowed {
log::warn!("Auth endpoint rate limit exceeded for IP: {}", ip);
if let Ok(audit_opt) = state.kit.require::<crate::registry::AuditModule>()
&& let Some(logger) = audit_opt.as_ref()
{
logger.log_rate_limit_exceeded(None, Some(ip.clone()));
}
if headers_enabled && let Some(values) = decision.headers {
let response = rate_limited_response();
return Ok(limiteron::middleware::inject_rate_limit_headers(
response, &values,
));
}
return Ok(rate_limited_response());
}
let response = next.run(request).await;
if headers_enabled && let Some(values) = decision.headers {
return Ok(limiteron::middleware::inject_rate_limit_headers(
response, &values,
));
}
Ok(response)
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::HeaderMap;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
fn peer_addr(ip: IpAddr) -> Option<SocketAddr> {
Some(SocketAddr::new(ip, 12345))
}
fn headers_with_xff(value: &str) -> HeaderMap {
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", value.parse().unwrap());
h
}
fn headers_with_x_real_ip(value: &str) -> HeaderMap {
let mut h = HeaderMap::new();
h.insert("x-real-ip", value.parse().unwrap());
h
}
#[test]
fn extract_ip_trusted_proxy_with_xff_returns_xff_ip() {
let peer = peer_addr(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)));
let headers = headers_with_xff("203.0.113.50");
let proxies = vec!["10.0.0.0/8".to_string()];
let result = extract_client_ip(&headers, peer, &proxies);
assert_eq!(result, Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 50))));
}
#[test]
fn extract_ip_untrusted_proxy_with_xff_returns_peer_ip() {
let peer = peer_addr(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100)));
let headers = headers_with_xff("203.0.113.50");
let proxies = vec!["10.0.0.0/8".to_string()];
let result = extract_client_ip(&headers, peer, &proxies);
assert_eq!(result, Some(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100))));
}
#[test]
fn extract_ip_trusted_proxy_xff_multiple_entries_returns_first() {
let peer = peer_addr(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)));
let headers = headers_with_xff("203.0.113.50, 70.41.32.12, 198.51.100.2");
let proxies = vec!["10.0.0.0/8".to_string()];
let result = extract_client_ip(&headers, peer, &proxies);
assert_eq!(result, Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 50))));
}
#[test]
fn extract_ip_empty_proxies_with_xff_returns_peer_ip() {
let peer = peer_addr(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100)));
let headers = headers_with_xff("203.0.113.50");
let proxies: Vec<String> = vec![];
let result = extract_client_ip(&headers, peer, &proxies);
assert_eq!(result, Some(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100))));
}
#[test]
fn requires_admin_matches_model_and_file_paths() {
assert!(requires_admin("/api/1/model/switch"));
assert!(requires_admin("/api/1/model/unload"));
assert!(requires_admin("/api/1/embed/file"));
assert!(requires_admin("/model/switch"));
assert!(requires_admin("/embed/file"));
assert!(!requires_admin("/api/1/embed"));
assert!(!requires_admin("/api/1/embed/batch"));
assert!(!requires_admin("/api/1/auth/login"));
assert!(!requires_admin("/health"));
}
#[test]
fn extract_ip_no_xff_returns_peer_ip() {
let peer = peer_addr(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100)));
let headers = HeaderMap::new();
let proxies = vec!["10.0.0.0/8".to_string()];
let result = extract_client_ip(&headers, peer, &proxies);
assert_eq!(result, Some(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100))));
}
#[test]
fn extract_ip_no_xff_no_connect_info_returns_none() {
let headers = HeaderMap::new();
let proxies: Vec<String> = vec![];
let result = extract_client_ip(&headers, None, &proxies);
assert_eq!(result, None);
}
#[test]
fn extract_ip_x_real_ip_used_when_xff_absent() {
let peer = peer_addr(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)));
let headers = headers_with_x_real_ip("203.0.113.99");
let proxies = vec!["10.0.0.0/8".to_string()];
let result = extract_client_ip(&headers, peer, &proxies);
assert_eq!(result, Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 99))));
}
#[test]
fn extract_ip_xff_preferred_over_x_real_ip() {
let peer = peer_addr(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)));
let mut headers = headers_with_xff("203.0.113.50");
headers.insert("x-real-ip", "198.51.100.7".parse().unwrap());
let proxies = vec!["10.0.0.0/8".to_string()];
let result = extract_client_ip(&headers, peer, &proxies);
assert_eq!(result, Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 50))));
}
#[test]
fn extract_ip_ipv6_peer_with_trusted_proxy() {
let peer = peer_addr(IpAddr::V6(Ipv6Addr::LOCALHOST));
let headers = headers_with_xff("2001:db8::1");
let proxies = vec!["10.0.0.0/8".to_string()];
let result = extract_client_ip(&headers, peer, &proxies);
assert_eq!(result, Some(IpAddr::V6(Ipv6Addr::LOCALHOST)));
}
#[test]
fn forbidden_response_has_correct_status_and_json_body() {
let resp = forbidden_response();
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
assert_eq!(
resp.headers()
.get("content-type")
.unwrap()
.to_str()
.unwrap(),
"application/json"
);
}
#[test]
fn unauthorized_response_has_structured_json_body() {
let resp = unauthorized_response("auth-credentials-missing");
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(
resp.headers()
.get("content-type")
.unwrap()
.to_str()
.unwrap(),
"application/json"
);
}
#[test]
fn rate_limited_response_has_structured_json_body() {
let resp = rate_limited_response();
assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(
resp.headers()
.get("content-type")
.unwrap()
.to_str()
.unwrap(),
"application/json"
);
}
#[test]
fn public_paths_do_not_require_admin() {
for path in PUBLIC_PATHS {
assert!(!requires_admin(path), "公开路径 {path} 不应要求 admin 角色");
}
}
#[test]
fn admin_path_prefixes_cover_all_dangerous_endpoints() {
assert!(requires_admin("/api/1/model/switch"));
assert!(requires_admin("/api/1/model/unload"));
assert!(requires_admin("/api/1/model/current"));
assert!(requires_admin("/api/1/model/info"));
assert!(requires_admin("/api/1/embed/file"));
assert!(requires_admin("/model/switch"));
assert!(requires_admin("/embed/file"));
}
#[test]
fn normal_endpoints_not_blocked_by_admin_check() {
let safe_paths = [
"/api/1/embed",
"/api/1/embed/batch",
"/api/1/similarity",
"/api/1/rerank",
"/api/1/auth/login",
"/api/1/auth/refresh",
"/api/1/auth/logout",
"/health",
"/metrics",
"/api-docs",
];
for path in safe_paths {
assert!(!requires_admin(path), "普通端点 {path} 不应要求 admin 角色");
}
}
}