use cedar_policy::{Context, RestrictedExpression};
use chrono::{DateTime, Utc};
use super::error::AuthzError;
pub trait BuildRequestContext: Send + Sync {
fn to_cedar_context(&self) -> Result<Context, AuthzError>;
}
pub struct NoContext;
impl BuildRequestContext for NoContext {
fn to_cedar_context(&self) -> Result<Context, AuthzError> {
Ok(Context::empty())
}
}
pub struct StandardRequestContext {
pub mfa_verified: bool,
pub ip_address: Option<std::net::IpAddr>,
pub timestamp: Option<DateTime<Utc>>,
}
impl StandardRequestContext {
pub fn new(mfa_verified: bool, ip_address: Option<std::net::IpAddr>) -> Self {
Self {
mfa_verified,
ip_address,
timestamp: None,
}
}
pub fn at(
mfa_verified: bool,
ip_address: Option<std::net::IpAddr>,
timestamp: DateTime<Utc>,
) -> Self {
Self {
mfa_verified,
ip_address,
timestamp: Some(timestamp),
}
}
}
impl BuildRequestContext for StandardRequestContext {
fn to_cedar_context(&self) -> Result<Context, AuthzError> {
let mut pairs: Vec<(String, RestrictedExpression)> = vec![(
"mfa_verified".to_string(),
RestrictedExpression::new_bool(self.mfa_verified),
)];
if let Some(ip) = &self.ip_address {
pairs.push((
"ip_address".to_string(),
RestrictedExpression::new_string(ip.to_string()),
));
}
if let Some(ts) = &self.timestamp {
pairs.push((
"timestamp".to_string(),
RestrictedExpression::new_string(ts.to_rfc3339()),
));
}
Context::from_pairs(pairs).map_err(|e| AuthzError::Context(format!("{e:?}")))
}
}
pub fn ip_from_headers_untrusted(headers: &axum::http::HeaderMap) -> Option<std::net::IpAddr> {
let raw = headers
.get("X-Real-IP")
.or_else(|| headers.get("X-Forwarded-For"))
.and_then(|v| v.to_str().ok())?;
raw.split(',').next().and_then(|s| s.trim().parse().ok())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum TrustedEntry {
Exact(std::net::IpAddr),
Net(std::net::IpAddr, u8),
}
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum CidrParseError {
#[error("missing '/' in CIDR: {0}")]
MissingPrefix(String),
#[error("invalid address in CIDR: {0}")]
Address(String),
#[error("invalid prefix length in CIDR: {0}")]
Prefix(String),
#[error("CIDR names a host, not a network (host bits set below the prefix): {0}")]
HostBitsSet(String),
}
fn parse_cidr(spec: &str) -> Result<(std::net::IpAddr, u8), CidrParseError> {
let (addr, prefix) = spec
.split_once('/')
.ok_or_else(|| CidrParseError::MissingPrefix(spec.to_string()))?;
let addr: std::net::IpAddr = addr
.trim()
.parse()
.map_err(|_| CidrParseError::Address(spec.to_string()))?;
let prefix: u8 = prefix
.trim()
.parse()
.map_err(|_| CidrParseError::Prefix(spec.to_string()))?;
let max = if addr.is_ipv4() { 32 } else { 128 };
if prefix > max {
return Err(CidrParseError::Prefix(spec.to_string()));
}
if !host_bits_clear(addr, prefix) {
return Err(CidrParseError::HostBitsSet(spec.to_string()));
}
Ok((addr, prefix))
}
fn host_bits_clear(addr: std::net::IpAddr, prefix: u8) -> bool {
match addr {
std::net::IpAddr::V4(a) => prefix == 32 || u32::from(a) & (u32::MAX >> prefix) == 0,
std::net::IpAddr::V6(a) => prefix == 128 || u128::from(a) & (u128::MAX >> prefix) == 0,
}
}
fn in_network(addr: std::net::IpAddr, net: std::net::IpAddr, prefix: u8) -> bool {
match (addr, net) {
(std::net::IpAddr::V4(a), std::net::IpAddr::V4(n)) => {
prefix == 0 || {
let shift = 32 - u32::from(prefix);
u32::from(a) >> shift == u32::from(n) >> shift
}
}
(std::net::IpAddr::V6(a), std::net::IpAddr::V6(n)) => {
prefix == 0 || {
let shift = 128 - u32::from(prefix);
u128::from(a) >> shift == u128::from(n) >> shift
}
}
_ => false,
}
}
#[derive(Debug, Clone, Default)]
pub struct TrustedProxies {
entries: std::sync::Arc<[TrustedEntry]>,
}
impl TrustedProxies {
pub fn new(addrs: impl IntoIterator<Item = std::net::IpAddr>) -> Self {
Self {
entries: addrs.into_iter().map(TrustedEntry::Exact).collect(),
}
}
pub fn from_cidrs<I, S>(specs: I) -> Result<Self, CidrParseError>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let entries = specs
.into_iter()
.map(|spec| parse_cidr(spec.as_ref()).map(|(net, p)| TrustedEntry::Net(net, p)))
.collect::<Result<Vec<_>, _>>()?;
Ok(Self {
entries: entries.into(),
})
}
pub fn with_cidrs<I, S>(self, specs: I) -> Result<Self, CidrParseError>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let mut entries: Vec<TrustedEntry> = self.entries.iter().copied().collect();
for spec in specs {
let (net, prefix) = parse_cidr(spec.as_ref())?;
entries.push(TrustedEntry::Net(net, prefix));
}
Ok(Self {
entries: entries.into(),
})
}
pub fn loopback_only() -> Self {
Self::new([
std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST),
std::net::IpAddr::V6(std::net::Ipv6Addr::LOCALHOST),
])
}
pub fn trusts(&self, peer: std::net::IpAddr) -> bool {
self.entries.iter().any(|entry| match *entry {
TrustedEntry::Exact(addr) => addr == peer,
TrustedEntry::Net(net, prefix) => in_network(peer, net, prefix),
})
}
}
pub fn ip_from_headers_trusted(
headers: &axum::http::HeaderMap,
peer: std::net::IpAddr,
trusted: &TrustedProxies,
) -> std::net::IpAddr {
if !trusted.trusts(peer) {
return peer;
}
let mut lines = headers.get_all("X-Forwarded-For").iter().peekable();
if lines.peek().is_some() {
let mut entries: Vec<&str> = Vec::new();
for value in lines {
let Ok(text) = value.to_str() else {
return peer;
};
entries.extend(text.split(','));
}
for entry in entries.iter().rev() {
match entry.trim().parse::<std::net::IpAddr>() {
Ok(addr) if trusted.trusts(addr) => continue,
Ok(addr) => return addr,
Err(_) => return peer,
}
}
return peer;
}
let mut real_ip = headers.get_all("X-Real-IP").iter();
match (real_ip.next(), real_ip.next()) {
(Some(value), None) => value
.to_str()
.ok()
.and_then(|value| value.trim().parse().ok())
.unwrap_or(peer),
_ => peer,
}
}
impl<T: BuildRequestContext> BuildRequestContext for std::sync::Arc<T> {
fn to_cedar_context(&self) -> Result<Context, AuthzError> {
self.as_ref().to_cedar_context()
}
}
#[cfg(test)]
mod tests;