use axum::http::HeaderMap;
use dashmap::DashMap;
use fraiseql_core::security::SecurityContext;
use fraiseql_error::{FraiseQLError, Result};
use tracing::warn;
pub(crate) const MAX_TENANT_KEY_LEN: usize = crate::tenancy::schema_isolation::MAX_PG_IDENTIFIER_LEN
- crate::tenancy::schema_isolation::TENANT_SCHEMA_PREFIX.len();
pub struct TenantKeyResolver;
impl TenantKeyResolver {
#[doc(hidden)] pub fn resolve(
security_context: Option<&SecurityContext>,
headers: &HeaderMap,
domain_registry: Option<&DomainRegistry>,
strict: bool,
) -> Result<Option<String>> {
let hints = Self::client_hints(headers, domain_registry)?;
if let Some(ctx) = security_context {
let bound = ctx.tenant_id.as_ref().map(|t| t.0.as_str());
if let Some((source, named)) =
hints.iter().find(|(_, named)| Some(named.as_str()) != bound)
{
warn!(
source,
named, "authenticated request names a tenant its token does not bind"
);
return Err(FraiseQLError::unauthorized(format!(
"the request names tenant '{named}' ({source}), which the caller's token \
is not bound to"
)));
}
return Ok(bound.map(str::to_string));
}
if let [(first, a), rest @ ..] = hints.as_slice() {
if let Some((second, b)) = rest.iter().find(|(_, b)| b != a) {
warn!("Tenant source conflict detected: {first}: {a}, {second}: {b}");
if strict {
return Err(FraiseQLError::Validation {
message: format!(
"Conflicting tenant values from sources: {first}: {a}, {second}: {b}"
),
path: None,
});
}
}
}
Ok(hints.into_iter().next().map(|(_, key)| key))
}
fn client_hints(
headers: &HeaderMap,
domain_registry: Option<&DomainRegistry>,
) -> Result<Vec<(&'static str, String)>> {
let mut hints = Vec::new();
if let Some(value) = headers.get("X-Tenant-ID").and_then(|v| v.to_str().ok()) {
validate_tenant_key(value)?;
hints.push(("X-Tenant-ID", value.to_string()));
}
if let (Some(registry), Some(host)) =
(domain_registry, headers.get("Host").and_then(|v| v.to_str().ok()))
{
if let Some(key) = registry.lookup(host) {
hints.push(("Host", key));
}
}
Ok(hints)
}
}
pub(crate) fn validate_tenant_key(key: &str) -> Result<()> {
if key.len() > MAX_TENANT_KEY_LEN {
return Err(FraiseQLError::validation(format!(
"X-Tenant-ID exceeds maximum length of {MAX_TENANT_KEY_LEN} characters"
)));
}
if !key.bytes().all(|b| b.is_ascii_alphanumeric() || b == b'_') {
return Err(FraiseQLError::validation(
"X-Tenant-ID contains invalid characters (allowed: a-zA-Z0-9_)",
));
}
Ok(())
}
pub struct DomainRegistry {
domains: DashMap<String, String>,
}
impl DomainRegistry {
#[must_use]
pub fn new() -> Self {
Self {
domains: DashMap::new(),
}
}
pub fn register(&self, domain: impl Into<String>, tenant_key: impl Into<String>) {
self.domains.insert(domain.into(), tenant_key.into());
}
#[must_use]
pub fn remove(&self, domain: &str) -> bool {
self.domains.remove(domain).is_some()
}
#[must_use]
pub fn lookup(&self, host: &str) -> Option<String> {
let domain = host.split(':').next().unwrap_or(host);
self.domains.get(domain).map(|v| v.clone())
}
#[must_use]
pub fn domains(&self) -> Vec<(String, String)> {
self.domains.iter().map(|e| (e.key().clone(), e.value().clone())).collect()
}
#[must_use]
pub fn len(&self) -> usize {
self.domains.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.domains.is_empty()
}
}
impl Default for DomainRegistry {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests;