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 tower_http::auth::AsyncAuthorizeRequest;
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 SurrealAuth;
impl AsyncAuthorizeRequest<Body> for SurrealAuth {
type RequestBody = Body;
type ResponseBody = Body;
type Future = BoxFuture<'static, Result<Request<Body>, Response<Self::ResponseBody>>>;
fn authorize(&mut self, request: Request<Body>) -> Self::Future {
Box::pin(async {
let (mut parts, body) = request.into_parts();
match check_auth(&mut parts).await {
Ok(sess) => {
parts.extensions.insert(sess);
Ok(Request::from_parts(parts, body))
}
Err(err) => {
let unauthorized_response = 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
});
Err(unauthorized_response)
}
}
})
}
}
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)
}