#[cfg(test)]
mod tests;
use std::{convert::Infallible, future::Future, net::SocketAddr};
use axum::{
extract::{ConnectInfo, FromRequestParts, rejection::ExtensionRejection},
http::request::Parts,
};
use fraiseql_core::security::SecurityContext;
use crate::middleware::{AuthUser, TenantClaim};
pub struct PeerIp(pub String);
impl<S> FromRequestParts<S> for PeerIp
where
S: Send + Sync,
{
type Rejection = Infallible;
fn from_request_parts(
parts: &mut Parts,
_state: &S,
) -> impl Future<Output = Result<Self, Self::Rejection>> + Send {
let ip = parts
.extensions
.get::<ConnectInfo<SocketAddr>>()
.map_or_else(|| "unknown".to_string(), |ci| ci.0.ip().to_string());
async move { Ok(PeerIp(ip)) }
}
}
#[derive(Debug, Clone)]
pub struct OptionalSecurityContext(pub Option<SecurityContext>);
impl<S> FromRequestParts<S> for OptionalSecurityContext
where
S: Send + Sync + 'static,
{
type Rejection = ExtensionRejection;
#[allow(clippy::manual_async_fn)] fn from_request_parts(
parts: &mut Parts,
_state: &S,
) -> impl Future<Output = Result<Self, Self::Rejection>> + Send {
async move {
let auth_user: Option<AuthUser> = parts.extensions.get::<AuthUser>().cloned();
let tenant_claim = parts.extensions.get::<TenantClaim>().cloned();
let headers = &parts.headers;
let security_context = auth_user.map(|auth_user| {
let request_id = extract_request_id(headers);
let mut context = build_security_context(
&auth_user.0,
request_id,
tenant_claim.as_ref().map(|c| &*c.0),
);
context.ip_address = extract_ip_address(headers);
context
});
Ok(OptionalSecurityContext(security_context))
}
}
}
#[must_use]
pub(crate) fn build_security_context(
user: &fraiseql_core::security::AuthenticatedUser,
request_id: String,
tenant_claim: Option<&str>,
) -> SecurityContext {
let mut context = SecurityContext::from_user(user, request_id);
for (key, value) in &user.extra_claims {
if key.starts_with("fraiseql.") {
continue;
}
context.attributes.insert(key.clone(), value.clone());
}
context.tenant_id = tenant_claim
.and_then(|claim| context.jwt_claim(claim, claim))
.and_then(|value| match value {
serde_json::Value::String(s) if !s.is_empty() => Some(s),
serde_json::Value::Number(n) => Some(n.to_string()),
_ => None,
})
.map(fraiseql_core::types::TenantId::new);
context
}
pub(crate) fn extract_request_id(headers: &axum::http::HeaderMap) -> String {
headers
.get("x-request-id")
.and_then(|v| v.to_str().ok())
.map_or_else(|| format!("req-{}", uuid::Uuid::new_v4()), |s| s.to_string())
}
pub(crate) const fn extract_ip_address(_headers: &axum::http::HeaderMap) -> Option<String> {
None
}