use std::task::{Context, Poll};
use anyhow::{Result, bail};
use axum::body::Body;
use axum::{Extension, RequestPartsExt};
use axum_extra::TypedHeader;
use axum_extra::headers::authorization::{Basic, Bearer};
use axum_extra::headers::{Authorization, Origin};
use futures_util::future::BoxFuture;
use http::StatusCode;
use http::request::Parts;
use hyper::{Request, Response};
use surrealdb_core::dbs::Session;
use surrealdb_core::iam::verify::{basic, token};
use surrealdb_core::observe::HttpRequestEventCtx;
use tower::{Layer, Service};
use uuid::Uuid;
use super::AppState;
use super::client_ip::ExtractClientIP;
use super::headers::{
SurrealAuthDatabase, SurrealAuthNamespace, SurrealDatabase, SurrealId, SurrealNamespace,
parse_typed_header,
};
use crate::ntw::error::Error as NetError;
#[derive(Clone, Copy)]
pub(super) struct SurrealAuthLayer;
impl<S> Layer<S> for SurrealAuthLayer {
type Service = SurrealAuthService<S>;
fn layer(&self, inner: S) -> Self::Service {
SurrealAuthService {
inner,
}
}
}
#[derive(Clone)]
pub(super) struct SurrealAuthService<S> {
inner: S,
}
impl<S> Service<Request<Body>> for SurrealAuthService<S>
where
S: Service<Request<Body>, Response = Response<Body>> + Clone + Send + 'static,
S::Future: Send + 'static,
S::Error: Send + 'static,
{
type Response = Response<Body>;
type Error = S::Error;
type Future = BoxFuture<'static, Result<Response<Body>, S::Error>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request<Body>) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
Box::pin(async move {
let (mut parts, body) = request.into_parts();
match check_auth(&mut parts).await {
Ok(sess) => {
let ctx = HttpRequestEventCtx::from_session(&sess);
parts.extensions.insert(sess);
let mut response = inner.call(Request::from_parts(parts, body)).await?;
response.extensions_mut().insert(ctx);
Ok(response)
}
Err(err) => {
let unauthorized = Response::builder()
.status(StatusCode::UNAUTHORIZED)
.body(Body::new(err.to_string()))
.unwrap_or_else(|_| {
let mut resp = Response::new(Body::empty());
*resp.status_mut() = StatusCode::UNAUTHORIZED;
resp
});
Ok(unauthorized)
}
}
})
}
}
async fn check_auth(parts: &mut Parts) -> Result<Session> {
let or = match parts.extract::<TypedHeader<Origin>>().await {
Ok(or) => {
if !or.is_null() {
Some(or.to_string())
} else {
None
}
}
_ => None,
};
let id = match parse_typed_header::<SurrealId>(parts.extract::<TypedHeader<SurrealId>>().await)?
{
Some(id) => {
match Uuid::try_parse(&id) {
Ok(id) => Some(id),
Err(_) => bail!(NetError::Request),
}
}
None => Some(Uuid::new_v4()),
};
let ns = parse_typed_header::<SurrealNamespace>(
parts.extract::<TypedHeader<SurrealNamespace>>().await,
)?;
let db = parse_typed_header::<SurrealDatabase>(
parts.extract::<TypedHeader<SurrealDatabase>>().await,
)?;
let auth_ns = parse_typed_header::<SurrealAuthNamespace>(
parts.extract::<TypedHeader<SurrealAuthNamespace>>().await,
)?;
let auth_db = parse_typed_header::<SurrealAuthDatabase>(
parts.extract::<TypedHeader<SurrealAuthDatabase>>().await,
)?;
let Extension(state) = parts.extract::<Extension<AppState>>().await.map_err(|err| {
tracing::error!("Error extracting the app state: {:?}", err);
NetError::InvalidAuth
})?;
let kvs = &state.datastore;
let ExtractClientIP(ip) =
parts.extract_with_state(&state).await.unwrap_or(ExtractClientIP(None));
let mut session = Session {
ip,
or,
id,
ns,
db,
..Session::default()
};
if let Ok(au) = parts.extract::<TypedHeader<Authorization<Basic>>>().await {
basic(
kvs,
&mut session,
au.username(),
au.password(),
auth_ns.as_deref(),
auth_db.as_deref(),
)
.await?;
};
if let Ok(au) = parts.extract::<TypedHeader<Authorization<Bearer>>>().await {
token(kvs, &mut session, au.token()).await?;
};
Ok(session)
}