use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use tonic::codegen::http::{Request, Response};
use tonic::codegen::Service;
use tonic::metadata::MetadataMap;
use tower::Layer;
use crate::grpc::session::{SessionTenant, TenantResolver, TenantScope};
#[derive(Clone)]
pub struct TenantResolverLayer {
resolver: Arc<dyn TenantResolver>,
}
impl TenantResolverLayer {
pub fn new(resolver: Arc<dyn TenantResolver>) -> Self {
Self { resolver }
}
}
impl<S> Layer<S> for TenantResolverLayer {
type Service = TenantResolved<S>;
fn layer(&self, inner: S) -> Self::Service {
TenantResolved {
inner,
resolver: Arc::clone(&self.resolver),
}
}
}
#[derive(Clone)]
pub struct TenantResolved<S> {
inner: S,
resolver: Arc<dyn TenantResolver>,
}
impl<S> tonic::server::NamedService for TenantResolved<S>
where
S: tonic::server::NamedService,
{
const NAME: &'static str = S::NAME;
}
impl<S, ReqBody, ResBody> Service<Request<ReqBody>> for TenantResolved<S>
where
S: Service<Request<ReqBody>, Response = Response<ResBody>> + Clone + Send + 'static,
S::Future: Send + 'static,
ReqBody: Send + 'static,
ResBody: Default,
{
type Response = S::Response;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: Request<ReqBody>) -> Self::Future {
let resolver = Arc::clone(&self.resolver);
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
Box::pin(async move {
let metadata = MetadataMap::from_headers(req.headers().clone());
match resolver.resolve(&metadata).await {
Ok(scope) => {
let tenant = match scope {
TenantScope::Tenant(t) => Some(t),
TenantScope::Global => None,
};
req.extensions_mut().insert(SessionTenant(tenant));
inner.call(req).await
}
Err(status) => Ok(status.into_http()),
}
})
}
}