use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use tonic::metadata::{MetadataMap, MetadataValue};
use tonic::Status;
use jammi_db::TenantId;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct SessionId(String);
impl SessionId {
pub fn new(s: impl Into<String>) -> Self {
Self(s.into())
}
pub fn as_str(&self) -> &str {
&self.0
}
}
#[derive(Debug, Clone, Copy)]
pub struct SessionTenant(pub Option<TenantId>);
#[derive(Debug, Default, Clone)]
pub struct SessionStore {
inner: Arc<RwLock<HashMap<SessionId, Option<TenantId>>>>,
}
impl SessionStore {
pub fn new() -> Self {
Self::default()
}
pub fn set(&self, session: SessionId, tenant: Option<TenantId>) {
self.inner
.write()
.expect("session store lock poisoned")
.insert(session, tenant);
}
pub fn get(&self, session: &SessionId) -> Option<TenantId> {
self.inner
.read()
.expect("session store lock poisoned")
.get(session)
.copied()
.flatten()
}
pub fn clear(&self, session: &SessionId) {
self.inner
.write()
.expect("session store lock poisoned")
.remove(session);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TenantScope {
Tenant(TenantId),
Global,
}
#[async_trait::async_trait]
pub trait TenantResolver: Send + Sync + 'static {
async fn resolve(&self, metadata: &MetadataMap) -> Result<TenantScope, Status>;
}
pub use jammi_wire::SESSION_HEADER;
#[derive(Clone)]
pub struct SessionIdTenantResolver {
store: SessionStore,
}
impl SessionIdTenantResolver {
pub fn new(store: SessionStore) -> Self {
Self { store }
}
pub fn arc(store: SessionStore) -> Arc<dyn TenantResolver> {
Arc::new(Self::new(store))
}
}
#[async_trait::async_trait]
impl TenantResolver for SessionIdTenantResolver {
async fn resolve(&self, metadata: &MetadataMap) -> Result<TenantScope, Status> {
let tenant = read_session_header(metadata).and_then(|sid| self.store.get(&sid));
Ok(match tenant {
Some(t) => TenantScope::Tenant(t),
None => TenantScope::Global,
})
}
}
fn read_session_header(metadata: &MetadataMap) -> Option<SessionId> {
metadata
.get(SESSION_HEADER)
.and_then(|v: &MetadataValue<_>| v.to_str().ok())
.map(SessionId::new)
}
#[cfg(test)]
mod tests {
use super::*;
use std::str::FromStr;
use tonic::Request;
fn t_a() -> TenantId {
TenantId::from_str("01906c83-d4c8-7e10-9c4f-3b6f7c5a8e9a").unwrap()
}
fn t_b() -> TenantId {
TenantId::from_str("01906c83-d4c8-7e10-9c4f-3b6f7c5a8e9b").unwrap()
}
#[test]
fn store_get_set_clear_roundtrip() {
let store = SessionStore::new();
let sid = SessionId::new("conn-1");
assert!(store.get(&sid).is_none());
store.set(sid.clone(), Some(t_a()));
assert_eq!(store.get(&sid), Some(t_a()));
store.set(sid.clone(), Some(t_b()));
assert_eq!(store.get(&sid), Some(t_b()));
store.set(sid.clone(), None);
assert!(store.get(&sid).is_none());
store.set(sid.clone(), Some(t_a()));
store.clear(&sid);
assert!(store.get(&sid).is_none());
}
#[test]
fn store_isolates_sessions() {
let store = SessionStore::new();
let s1 = SessionId::new("conn-1");
let s2 = SessionId::new("conn-2");
store.set(s1.clone(), Some(t_a()));
store.set(s2.clone(), Some(t_b()));
assert_eq!(store.get(&s1), Some(t_a()));
assert_eq!(store.get(&s2), Some(t_b()));
}
#[tokio::test]
async fn session_id_resolver_is_global_when_no_header() {
let resolver = SessionIdTenantResolver::new(SessionStore::new());
let scope = resolver.resolve(Request::new(()).metadata()).await.unwrap();
assert_eq!(
scope,
TenantScope::Global,
"no session header is the explicit Global scope, not an error"
);
}
#[tokio::test]
async fn session_id_resolver_binds_tenant_when_header_present() {
let store = SessionStore::new();
store.set(SessionId::new("conn-1"), Some(t_a()));
let resolver = SessionIdTenantResolver::new(store);
let mut req = Request::new(());
req.metadata_mut()
.insert(SESSION_HEADER, "conn-1".parse().unwrap());
let scope = resolver.resolve(req.metadata()).await.unwrap();
assert_eq!(scope, TenantScope::Tenant(t_a()));
}
#[tokio::test]
async fn session_id_resolver_is_global_when_session_unbound() {
let resolver = SessionIdTenantResolver::new(SessionStore::new());
let mut req = Request::new(());
req.metadata_mut()
.insert(SESSION_HEADER, "unknown".parse().unwrap());
let scope = resolver.resolve(req.metadata()).await.unwrap();
assert_eq!(scope, TenantScope::Global);
}
}