use std::{collections::HashMap, sync::Arc};
use axum::http::{HeaderMap, HeaderName};
use fraiseql_core::security::{ENRICHED_NAMESPACE_PREFIX, SecurityContext};
use serde::Deserialize;
use subtle::ConstantTimeEq;
use tracing::{debug, warn};
use crate::api_key::sha256_hash;
const SA_HEADER: &str = "x-api-key";
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ServiceAccountConfig {
pub secret_env: String,
#[serde(default)]
pub roles: Vec<String>,
#[serde(default)]
pub scopes: Vec<String>,
#[serde(default)]
pub tenant: Option<String>,
#[serde(default)]
pub static_enriched: HashMap<String, serde_json::Value>,
}
#[derive(Debug)]
#[non_exhaustive]
pub enum SaAuth {
NoSecret,
Ambiguous,
Authenticated(Box<SecurityContext>),
Unmatched,
}
#[derive(Debug, Clone)]
struct ResolvedServiceAccount {
name: String,
secret_hash: [u8; 32],
roles: Vec<String>,
scopes: Vec<String>,
tenant: Option<String>,
static_enriched: HashMap<String, serde_json::Value>,
}
pub struct ServiceAccountAuthenticator {
header_name: HeaderName,
accounts: Vec<ResolvedServiceAccount>,
}
impl std::fmt::Debug for ServiceAccountAuthenticator {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ServiceAccountAuthenticator")
.field("header_name", &self.header_name)
.field("accounts_count", &self.accounts.len())
.finish()
}
}
impl ServiceAccountAuthenticator {
#[must_use]
pub fn from_config(
accounts: &HashMap<String, ServiceAccountConfig>,
resolve_secret: impl Fn(&str) -> Option<String>,
) -> Option<Arc<Self>> {
let mut resolved = Vec::new();
for (name, cfg) in accounts {
match resolve_secret(&cfg.secret_env) {
Some(secret) if !secret.is_empty() => {
resolved.push(ResolvedServiceAccount {
name: name.clone(),
secret_hash: sha256_hash(secret.as_bytes()),
roles: cfg.roles.clone(),
scopes: cfg.scopes.clone(),
tenant: cfg.tenant.clone(),
static_enriched: cfg.static_enriched.clone(),
});
},
_ => warn!(
account = %name,
secret_env = %cfg.secret_env,
"service account skipped — its secret env var is unset or empty"
),
}
}
if resolved.is_empty() {
return None;
}
let header_name = HeaderName::from_static(SA_HEADER);
Some(Arc::new(Self {
header_name,
accounts: resolved,
}))
}
#[must_use]
pub fn resolve(&self, headers: &HeaderMap, jwt_present: bool) -> SaAuth {
if !self.header_present(headers) {
return SaAuth::NoSecret;
}
if jwt_present {
return SaAuth::Ambiguous;
}
self.authenticate(headers).map_or(SaAuth::Unmatched, SaAuth::Authenticated)
}
#[must_use]
pub fn header_present(&self, headers: &HeaderMap) -> bool {
headers
.get(&self.header_name)
.and_then(|v| v.to_str().ok())
.is_some_and(|s| !s.is_empty())
}
#[must_use]
pub fn authenticate(&self, headers: &HeaderMap) -> Option<Box<SecurityContext>> {
let raw = headers.get(&self.header_name)?.to_str().ok()?;
if raw.is_empty() {
return None;
}
let secret = strip_scheme_prefix(raw);
let presented = sha256_hash(secret.as_bytes());
for account in &self.accounts {
if bool::from(presented.ct_eq(&account.secret_hash)) {
debug!(account = %account.name, "service account authenticated");
return Some(Box::new(build_context(account)));
}
}
warn!("service-account authentication failed: no matching account");
None
}
}
fn strip_scheme_prefix(raw: &str) -> &str {
let bytes = raw.as_bytes();
if bytes.len() > 7
&& (bytes[..7].eq_ignore_ascii_case(b"apikey ")
|| bytes[..7].eq_ignore_ascii_case(b"bearer "))
{
&raw[7..]
} else {
raw
}
}
fn build_context(account: &ResolvedServiceAccount) -> SecurityContext {
let tenant = account.tenant.as_ref().map(fraiseql_core::types::TenantId::new);
let request_id = format!("sa-{}", uuid::Uuid::new_v4());
let mut ctx = SecurityContext::service_account(
account.name.clone(),
request_id,
account.roles.clone(),
account.scopes.clone(),
tenant,
);
for (field, value) in &account.static_enriched {
ctx.attributes
.insert(format!("{ENRICHED_NAMESPACE_PREFIX}{field}"), value.clone());
}
ctx
}
#[must_use]
pub fn service_account_authenticator_from_schema(
schema: &fraiseql_core::schema::CompiledSchema,
) -> Option<Arc<ServiceAccountAuthenticator>> {
let security = schema.security.as_ref()?;
let value = security.additional.get("service_accounts")?;
let accounts: HashMap<String, ServiceAccountConfig> = serde_json::from_value(value.clone())
.map_err(|e| warn!(error = %e, "Failed to parse security.service_accounts config"))
.ok()?;
ServiceAccountAuthenticator::from_config(&accounts, |env| std::env::var(env).ok())
}
#[cfg(test)]
mod tests;