use core::fmt;
use std::sync::Arc;
use connectrpc::interceptor::{
NextStream, PayloadStream, StreamRequest, StreamResponse, UnaryRequest, UnaryResponse,
};
use connectrpc::{ConnectError, Interceptor, Next, RequestContext, async_trait};
use futures::StreamExt as _;
use super::Shared;
use crate::auth_v2;
use crate::error::ServerError;
use crate::op::Procedure;
use crate::pipeline::{HookSet, Pipeline, RequestMeta};
use crate::principal::Principal;
use crate::store::{MultipartBlobStore, NamespaceStore};
pub struct AuthInterceptor<B, N, H> {
pipe: Shared<Pipeline<B, N, H>>,
}
impl<B, N, H> fmt::Debug for AuthInterceptor<B, N, H> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AuthInterceptor").finish_non_exhaustive()
}
}
impl<B: MultipartBlobStore, N: NamespaceStore, H: HookSet> AuthInterceptor<B, N, H> {
#[must_use]
pub fn new(pipeline: Arc<Pipeline<B, N, H>>) -> Self {
Self {
pipe: Shared::new(pipeline),
}
}
fn header(ctx: &RequestContext) -> impl Fn(&str) -> Option<String> + '_ {
move |name: &str| {
if name == "x-write-grant" {
let values: Vec<_> = ctx.headers().get_all(name).iter().collect();
if values.is_empty() {
return None;
}
return Some(
values
.iter()
.map(|value| value.to_str().unwrap_or("~"))
.collect::<Vec<_>>()
.join(", "),
);
}
ctx.header(name).and_then(|v| {
if name == "x-mkit-ref" {
core::str::from_utf8(v.as_bytes()).ok().map(str::to_owned)
} else if auth_v2::HEADER_NAMES.contains(&name) {
Some(v.to_str().unwrap_or("~").to_owned())
} else {
v.to_str().ok().map(str::to_owned)
}
})
}
}
fn authenticate(
&self,
ctx: &mut RequestContext,
unary_body: Option<&[u8]>,
) -> Result<(), ServerError> {
let Some(procedure) = ctx.path().and_then(Procedure::from_connect_path) else {
return Ok(());
};
let transport_principal = ctx.extensions().get::<Principal>().cloned();
let authenticated = {
let header = Self::header(ctx);
let meta = RequestMeta {
procedure,
header: &header,
header_values: Some(&|name| {
ctx.headers()
.get_all(name)
.iter()
.map(|value| value.to_str().unwrap_or("\n").to_owned())
.collect()
}),
unary_body,
transport_principal,
};
self.pipe.get().authenticate(&meta)?
};
ctx.extensions_mut().insert(authenticated);
Ok(())
}
}
#[async_trait]
impl<B, N, H> Interceptor for AuthInterceptor<B, N, H>
where
B: MultipartBlobStore + 'static,
N: NamespaceStore + 'static,
H: HookSet + 'static,
{
async fn intercept_unary(
&self,
mut req: UnaryRequest,
next: Next<'_>,
) -> Result<UnaryResponse, ConnectError> {
let body = req.payload.encoded()?;
self.authenticate(&mut req.ctx, Some(&body))?;
next.run(req).await
}
async fn intercept_streaming(
&self,
mut req: StreamRequest,
mut inbound: PayloadStream,
next: NextStream<'_>,
) -> Result<StreamResponse, ConnectError> {
if req.ctx.path().and_then(Procedure::from_connect_path) != Some(Procedure::DownloadPack) {
self.authenticate(&mut req.ctx, None)?;
return next.run(req, inbound).await;
}
let Some(item) = inbound.next().await else {
self.authenticate(&mut req.ctx, None)?;
return next.run(req, inbound).await;
};
let payload = item?;
let bytes = payload.encoded()?;
if bytes.len() > 1024 {
return Err(ServerError::invalid_argument("DownloadPack request too large").into());
}
let compressed_signed = {
let header = Self::header(&req.ctx);
auth_v2::carries_auth_headers(&header)
&& ["connect-content-encoding", "grpc-encoding"]
.iter()
.any(|name| {
req.ctx
.header(*name)
.is_some_and(|v| !matches!(v.to_str(), Ok("identity")))
})
};
if compressed_signed {
return Err(ServerError::unauthenticated("compressed signed request").into());
}
let mut frame = Vec::with_capacity(5 + bytes.len());
frame.push(0);
frame.extend_from_slice(
&u32::try_from(bytes.len())
.map_err(|_| ServerError::invalid_argument("DownloadPack request too large"))?
.to_be_bytes(),
);
frame.extend_from_slice(&bytes);
self.authenticate(&mut req.ctx, Some(&frame))?;
let inbound: PayloadStream =
Box::pin(futures::stream::once(async move { Ok(payload) }).chain(inbound));
next.run(req, inbound).await
}
}