use super::{GatewayHaTenantBinding, GatewayRegistryError, GatewayRegistryResult};
use crate::config::MAX_GATEWAY_CONNECTIONS;
use crate::GatewayWorkerRegistration;
use appcore_types::{ClusterId, TenantId};
use std::collections::HashSet;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GatewayHaWorkerSnapshot {
pub tenant_id: TenantId,
pub cluster_id: ClusterId,
pub registration: GatewayWorkerRegistration,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GatewayHaSessionSnapshot {
pub tenant_id: TenantId,
pub session_id: String,
pub expires_at_ms: u64,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct GatewayHaOwnershipSnapshot {
pub workers: Vec<GatewayHaWorkerSnapshot>,
pub sessions: Vec<GatewayHaSessionSnapshot>,
}
impl GatewayHaOwnershipSnapshot {
pub fn validate(
&self,
tenants: &[GatewayHaTenantBinding],
now_ms: u64,
) -> GatewayRegistryResult<()> {
if self.workers.len().saturating_add(self.sessions.len()) > MAX_GATEWAY_CONNECTIONS {
return Err(GatewayRegistryError::CapacityExceeded);
}
let configured = tenants
.iter()
.map(|binding| binding.tenant_id.as_str())
.collect::<HashSet<_>>();
let mut workers = HashSet::with_capacity(self.workers.len());
for worker in &self.workers {
worker.registration.validate()?;
let configured_cluster = tenants
.iter()
.find(|binding| binding.tenant_id == worker.tenant_id)
.map(|binding| &binding.cluster_id);
if configured_cluster != Some(&worker.cluster_id)
|| !workers.insert(format!(
"{}\0{}\0{}",
worker.tenant_id.as_str(),
worker.registration.installation_id.as_str(),
worker.registration.core_id.as_str()
))
{
return Err(GatewayRegistryError::InvalidContract);
}
}
let mut sessions = HashSet::with_capacity(self.sessions.len());
for session in &self.sessions {
if !configured.contains(session.tenant_id.as_str())
|| session.expires_at_ms <= now_ms
|| !valid_identifier(&session.session_id)
|| !sessions.insert((session.tenant_id.as_str(), session.session_id.as_str()))
{
return Err(GatewayRegistryError::InvalidContract);
}
}
Ok(())
}
}
pub trait GatewayHaOwnershipSource: Send + Sync {
fn snapshot(&self, now_ms: u64) -> GatewayRegistryResult<GatewayHaOwnershipSnapshot>;
}
impl GatewayHaOwnershipSource for GatewayHaOwnershipSnapshot {
fn snapshot(&self, _now_ms: u64) -> GatewayRegistryResult<GatewayHaOwnershipSnapshot> {
Ok(self.clone())
}
}
fn valid_identifier(value: &str) -> bool {
!value.is_empty()
&& value.len() <= 128
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b':'))
}
#[cfg(test)]
#[path = "ownership_tests.rs"]
mod tests;