use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use super::domain_types::{RoleName, Scope};
#[cfg(test)]
mod tests;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "source", content = "claim", rename_all = "snake_case")]
#[non_exhaustive]
pub enum InjectedParamSource {
Jwt(String),
Enrichment(String),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct RoleDefinition {
pub name: RoleName,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
pub scopes: Vec<Scope>,
}
impl RoleDefinition {
#[must_use]
pub fn new(name: impl Into<String>, scopes: Vec<String>) -> Self {
Self {
name: RoleName::new(name),
description: None,
scopes: scopes.into_iter().map(Scope::new).collect(),
}
}
#[must_use]
pub fn with_description(mut self, description: String) -> Self {
self.description = Some(description);
self
}
#[must_use]
pub fn has_scope(&self, required_scope: &str) -> bool {
self.scopes.iter().any(|scope| {
let scope = scope.as_str();
if scope == "*" {
return true; }
if scope == required_scope {
return true; }
if let Some(prefix) = scope.strip_suffix(":*") {
return required_scope
.strip_prefix(prefix)
.is_some_and(|rest| rest.starts_with(':'));
}
if let Some(prefix) = scope.strip_suffix('*') {
return (prefix.ends_with('.') || prefix.ends_with(':'))
&& required_scope.starts_with(prefix);
}
false
})
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum TenancyMode {
#[default]
None,
Row,
Schema,
}
impl std::fmt::Display for TenancyMode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::None => write!(f, "none"),
Self::Row => write!(f, "row"),
Self::Schema => write!(f, "schema"),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TenancyConfig {
#[serde(default)]
pub mode: TenancyMode,
#[serde(default = "default_tenant_claim")]
pub tenant_claim: String,
}
fn default_tenant_claim() -> String {
"tenant_id".to_string()
}
impl Default for TenancyConfig {
fn default() -> Self {
Self {
mode: TenancyMode::None,
tenant_claim: default_tenant_claim(),
}
}
}
fn is_default_tenancy(t: &TenancyConfig) -> bool {
*t == TenancyConfig::default()
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct RlsConfig {
pub enabled: bool,
}
fn is_default_rls(r: &RlsConfig) -> bool {
*r == RlsConfig::default()
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct SecurityConfig {
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub role_definitions: Vec<RoleDefinition>,
#[serde(skip_serializing_if = "Option::is_none")]
pub default_role: Option<String>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub multi_tenant: bool,
#[serde(default, skip_serializing_if = "is_default_rls")]
pub rls: RlsConfig,
#[serde(default, skip_serializing_if = "is_default_tenancy")]
pub tenancy: TenancyConfig,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cost_budget: Option<CostBudgetConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default_policy: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub rules: Vec<super::config_types::AuthorizationRule>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub policies: Vec<super::config_types::AuthorizationPolicy>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub field_auth: Vec<super::config_types::FieldAuthRule>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub enterprise: Option<super::config_types::EnterpriseSecurityConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error_sanitization: Option<ErrorSanitizationConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub rate_limiting: Option<RateLimitingSecurityConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub state_encryption: Option<StateEncryptionConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub pkce: Option<PkceSecurityConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub api_keys: Option<ApiKeySecurityConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub token_revocation: Option<TokenRevocationSecurityConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub trusted_documents: Option<TrustedDocumentsConfig>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub persisted_queries_only: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub service_accounts: Option<HashMap<String, ServiceAccountConfig>>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct ErrorSanitizationConfig {
pub enabled: bool,
pub hide_implementation_details: bool,
pub sanitize_database_errors: bool,
pub custom_error_message: Option<String>,
}
impl Default for ErrorSanitizationConfig {
fn default() -> Self {
Self {
enabled: false,
hide_implementation_details: true,
sanitize_database_errors: true,
custom_error_message: None,
}
}
}
pub const DEFAULT_RATE_LIMIT_MAX_BUCKETS: usize = 100_000;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct RateLimitingSecurityConfig {
pub enabled: bool,
pub requests_per_second: u32,
pub burst_size: u32,
pub auth_start_max_requests: u32,
pub auth_start_window_secs: u64,
pub auth_callback_max_requests: u32,
pub auth_callback_window_secs: u64,
pub auth_refresh_max_requests: u32,
pub auth_refresh_window_secs: u64,
pub auth_logout_max_requests: u32,
pub auth_logout_window_secs: u64,
pub failed_login_max_attempts: u32,
pub failed_login_lockout_secs: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub requests_per_second_per_user: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub redis_url: Option<String>,
pub trust_proxy_headers: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub trusted_proxy_cidrs: Option<Vec<String>>,
pub max_buckets: usize,
}
impl Default for RateLimitingSecurityConfig {
fn default() -> Self {
Self {
enabled: false,
requests_per_second: 100,
requests_per_second_per_user: None,
burst_size: 200,
auth_start_max_requests: 5,
auth_start_window_secs: 60,
auth_callback_max_requests: 10,
auth_callback_window_secs: 60,
auth_refresh_max_requests: 20,
auth_refresh_window_secs: 300,
auth_logout_max_requests: 30,
auth_logout_window_secs: 60,
failed_login_max_attempts: 10,
failed_login_lockout_secs: 900,
redis_url: None,
trust_proxy_headers: false,
trusted_proxy_cidrs: None,
max_buckets: DEFAULT_RATE_LIMIT_MAX_BUCKETS,
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum EncryptionAlgorithm {
#[default]
#[serde(rename = "chacha20-poly1305")]
Chacha20Poly1305,
#[serde(rename = "aes-256-gcm")]
Aes256Gcm,
}
impl std::fmt::Display for EncryptionAlgorithm {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Chacha20Poly1305 => f.write_str("chacha20-poly1305"),
Self::Aes256Gcm => f.write_str("aes-256-gcm"),
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum KeySource {
#[default]
Env,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct StateEncryptionConfig {
pub enabled: bool,
pub algorithm: EncryptionAlgorithm,
pub key_source: KeySource,
pub key_env: Option<String>,
}
impl Default for StateEncryptionConfig {
fn default() -> Self {
Self {
enabled: false,
algorithm: EncryptionAlgorithm::default(),
key_source: KeySource::Env,
key_env: Some("STATE_ENCRYPTION_KEY".to_string()),
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum CodeChallengeMethod {
#[default]
#[serde(rename = "S256")]
S256,
#[serde(rename = "plain")]
Plain,
}
impl CodeChallengeMethod {
#[must_use]
pub const fn as_str(&self) -> &'static str {
match self {
Self::S256 => "S256",
Self::Plain => "plain",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct PkceSecurityConfig {
pub enabled: bool,
pub code_challenge_method: CodeChallengeMethod,
pub state_ttl_secs: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub redis_url: Option<String>,
}
impl Default for PkceSecurityConfig {
fn default() -> Self {
Self {
enabled: false,
code_challenge_method: CodeChallengeMethod::S256,
state_ttl_secs: 600,
redis_url: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct ApiKeySecurityConfig {
pub enabled: bool,
pub header: String,
pub hash_algorithm: String,
pub storage: String,
#[serde(rename = "static")]
pub static_keys: Vec<StaticApiKeyEntry>,
}
impl Default for ApiKeySecurityConfig {
fn default() -> Self {
Self {
enabled: false,
header: "X-API-Key".to_string(),
hash_algorithm: "sha256".to_string(),
storage: "env".to_string(),
static_keys: vec![],
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct StaticApiKeyEntry {
pub key_hash: String,
#[serde(default)]
pub scopes: Vec<String>,
pub name: String,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum TrustedDocumentMode {
Strict,
#[default]
Permissive,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct TrustedDocumentsConfig {
pub enabled: bool,
pub mode: TrustedDocumentMode,
#[serde(skip_serializing_if = "Option::is_none")]
pub manifest_path: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub manifest_url: Option<String>,
pub reload_interval_secs: u64,
}
impl Default for TrustedDocumentsConfig {
fn default() -> Self {
Self {
enabled: false,
mode: TrustedDocumentMode::Permissive,
manifest_path: None,
manifest_url: None,
reload_interval_secs: 0,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct TokenRevocationSecurityConfig {
pub enabled: bool,
pub backend: String,
pub require_jti: bool,
pub fail_open: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub redis_url: Option<String>,
pub revoke_all_ttl_secs: u64,
}
impl Default for TokenRevocationSecurityConfig {
fn default() -> Self {
Self {
enabled: false,
backend: "memory".to_string(),
require_jti: true,
fail_open: false,
redis_url: None,
revoke_all_ttl_secs: 86_400,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, 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, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(deny_unknown_fields)]
pub struct CostBudgetConfig {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub per_request_max: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub per_tenant_per_minute_default: Option<u64>,
}
impl SecurityConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn add_role(&mut self, role: RoleDefinition) {
self.role_definitions.push(role);
}
#[must_use]
pub fn find_role(&self, name: &str) -> Option<&RoleDefinition> {
self.role_definitions.iter().find(|r| r.name == name)
}
#[must_use]
pub fn get_role_scopes(&self, role_name: &str) -> Vec<String> {
self.find_role(role_name)
.map(|role| role.scopes.iter().map(|s| s.to_string()).collect::<Vec<String>>())
.unwrap_or_default()
}
#[must_use]
pub fn role_has_scope(&self, role_name: &str, scope: &str) -> bool {
self.find_role(role_name).is_some_and(|role| role.has_scope(scope))
}
}