use super::error::ApiResponseError;
use super::serve::AppState;
use crate::config::ServerConfig;
use axum::extract::rejection::PathRejection;
use axum::extract::{FromRequest, FromRequestParts, Path as AxumPath, Query};
use axum::http::request::Parts;
use axum::http::{HeaderMap, StatusCode};
use axum::Json;
use futures::StreamExt;
use loonfs::{ByteStream, ErrorCode};
use loonfs_api::NamespaceId;
use std::convert::Infallible;
use std::sync::{Arc, Mutex};
use tokio::sync::OwnedSemaphorePermit;
pub(super) fn server_busy_error(what: &str) -> ApiResponseError {
ApiResponseError::new(
StatusCode::SERVICE_UNAVAILABLE,
ErrorCode::ServerBusy,
&format!("the server is at its concurrency limit for {what}; retry shortly"),
)
.with_retry_after(1)
}
pub(super) fn authorize(
config: &ServerConfig,
headers: &HeaderMap,
) -> Result<(), ApiResponseError> {
let Some(expected) = &config.auth_token else {
return Ok(());
};
let actual = headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.unwrap_or_default();
let expected = format!("Bearer {}", expected.expose());
if constant_time_eq(actual.as_bytes(), expected.as_bytes()) {
Ok(())
} else {
Err(ApiResponseError::new(
StatusCode::UNAUTHORIZED,
ErrorCode::Unauthorized,
"missing or invalid bearer token",
))
}
}
fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
}
fn parse_namespace_id(value: String) -> Result<NamespaceId, ApiResponseError> {
NamespaceId::parse(&value).map_err(ApiResponseError::invalid_namespace_id)
}
#[derive(Debug, serde::Deserialize)]
struct NamespaceSegment {
namespace: String,
}
pub(super) struct NamespaceIdPath(Result<NamespaceId, ApiResponseError>);
impl NamespaceIdPath {
pub(super) fn into_id(self) -> Result<NamespaceId, ApiResponseError> {
self.0
}
}
impl<S> FromRequestParts<S> for NamespaceIdPath
where
S: Send + Sync,
{
type Rejection = Infallible;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
match AxumPath::<NamespaceSegment>::from_request_parts(parts, state).await {
Ok(AxumPath(NamespaceSegment { namespace })) => Ok(Self(parse_namespace_id(namespace))),
Err(rejection) => Ok(Self(Err(invalid_path_params(&rejection)))),
}
}
}
fn invalid_path_params(rejection: &PathRejection) -> ApiResponseError {
ApiResponseError::new(
StatusCode::BAD_REQUEST,
ErrorCode::InvalidRequest,
&format!("invalid path parameters: {rejection}"),
)
}
pub(super) struct AppPath<T>(Result<T, ApiResponseError>);
impl<T> AppPath<T> {
pub(super) fn into_params(self) -> Result<T, ApiResponseError> {
self.0
}
}
impl<S, T> FromRequestParts<S> for AppPath<T>
where
S: Send + Sync,
T: serde::de::DeserializeOwned + Send,
{
type Rejection = Infallible;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
match AxumPath::<T>::from_request_parts(parts, state).await {
Ok(AxumPath(value)) => Ok(Self(Ok(value))),
Err(rejection) => Ok(Self(Err(invalid_path_params(&rejection)))),
}
}
}
pub(super) struct AppQuery<T>(Result<T, ApiResponseError>);
impl<T> AppQuery<T> {
pub(super) fn into_params(self) -> Result<T, ApiResponseError> {
self.0
}
}
impl<S, T> FromRequestParts<S> for AppQuery<T>
where
S: Send + Sync,
T: serde::de::DeserializeOwned,
{
type Rejection = Infallible;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
match Query::<T>::from_request_parts(parts, state).await {
Ok(Query(value)) => Ok(Self(Ok(value))),
Err(rejection) => Ok(Self(Err(ApiResponseError::new(
StatusCode::BAD_REQUEST,
ErrorCode::InvalidRequest,
&format!("invalid query parameters: {rejection}"),
)))),
}
}
}
pub(super) struct AppJson<T>(pub(super) T);
async fn extract_json<S, T>(
req: axum::extract::Request,
state: &S,
body_too_large: fn() -> ApiResponseError,
) -> Result<T, ApiResponseError>
where
T: serde::de::DeserializeOwned,
S: Send + Sync,
{
match Json::<T>::from_request(req, state).await {
Ok(Json(value)) => Ok(value),
Err(rejection) if rejection.status() == StatusCode::PAYLOAD_TOO_LARGE => {
Err(body_too_large())
}
Err(rejection) => Err(ApiResponseError::new(
StatusCode::BAD_REQUEST,
ErrorCode::InvalidRequest,
&rejection.body_text(),
)),
}
}
impl<T> FromRequest<AppState> for AppJson<T>
where
T: serde::de::DeserializeOwned,
{
type Rejection = ApiResponseError;
async fn from_request(
req: axum::extract::Request,
state: &AppState,
) -> Result<Self, Self::Rejection> {
authorize(&state.config, req.headers())?;
extract_json(req, state, json_body_too_large_error)
.await
.map(AppJson)
}
}
pub(super) struct UploadBodyStream {
body: axum::body::Body,
max_bytes: u64,
abort: Arc<Mutex<Option<UploadStreamAbort>>>,
_permit: OwnedSemaphorePermit,
}
enum UploadStreamAbort {
TooLarge,
Unreadable(String),
}
impl UploadBodyStream {
pub(super) fn into_stream(self) -> (ByteStream, UploadStreamOutcome) {
let Self {
body,
max_bytes,
abort,
_permit,
} = self;
let outcome = UploadStreamOutcome {
abort: Arc::clone(&abort),
_permit,
};
let mut read_bytes = 0u64;
let stream = body
.into_data_stream()
.map(move |chunk| match chunk {
Ok(chunk) => {
read_bytes += chunk.len() as u64;
if read_bytes > max_bytes {
return Err(record_abort(&abort, UploadStreamAbort::TooLarge));
}
Ok(chunk)
}
Err(error) => Err(record_abort(
&abort,
UploadStreamAbort::Unreadable(error.to_string()),
)),
})
.boxed();
(stream, outcome)
}
}
fn record_abort(
abort: &Arc<Mutex<Option<UploadStreamAbort>>>,
reason: UploadStreamAbort,
) -> loonfs::ObjectStoreError {
let message = match &reason {
UploadStreamAbort::TooLarge => "upload body exceeded this deployment's limit".to_owned(),
UploadStreamAbort::Unreadable(error) => format!("upload body unreadable: {error}"),
};
*abort.lock().unwrap_or_else(|err| err.into_inner()) = Some(reason);
loonfs::ObjectStoreError::transport("upload body", message)
}
pub(super) struct UploadStreamOutcome {
abort: Arc<Mutex<Option<UploadStreamAbort>>>,
_permit: OwnedSemaphorePermit,
}
impl UploadStreamOutcome {
pub(super) fn into_rejection(self) -> Option<ApiResponseError> {
match self
.abort
.lock()
.unwrap_or_else(|err| err.into_inner())
.take()
{
Some(UploadStreamAbort::TooLarge) => Some(upload_body_too_large_error()),
Some(UploadStreamAbort::Unreadable(error)) => Some(ApiResponseError::new(
StatusCode::BAD_REQUEST,
ErrorCode::InvalidRequest,
&format!("request body unreadable: {error}"),
)),
None => None,
}
}
}
impl FromRequest<AppState> for UploadBodyStream {
type Rejection = ApiResponseError;
async fn from_request(
req: axum::extract::Request,
state: &AppState,
) -> Result<Self, Self::Rejection> {
authorize(&state.config, req.headers())?;
let permit = state
.upload_permits
.clone()
.try_acquire_owned()
.map_err(|_| {
state.metrics.upload_rejected_as_busy();
server_busy_error("proxied uploads")
})?;
let max_bytes = state.config.max_upload_bytes;
if declared_content_length(req.headers()).is_some_and(|length| length > max_bytes) {
return Err(upload_body_too_large_error());
}
Ok(UploadBodyStream {
body: req.into_body(),
max_bytes,
abort: Arc::new(Mutex::new(None)),
_permit: permit,
})
}
}
fn declared_content_length(headers: &HeaderMap) -> Option<u64> {
headers
.get(axum::http::header::CONTENT_LENGTH)?
.to_str()
.ok()?
.parse()
.ok()
}
fn upload_body_too_large_error() -> ApiResponseError {
ApiResponseError::new(
StatusCode::PAYLOAD_TOO_LARGE,
ErrorCode::ContentTooLarge,
"request body exceeds this deployment's limit; check the \
`upload.max_content_bytes` capability limit, and use `direct_put` \
for large content when `core.uploads.direct_put` is advertised",
)
}
fn json_body_too_large_error() -> ApiResponseError {
ApiResponseError::new(
StatusCode::PAYLOAD_TOO_LARGE,
ErrorCode::ContentTooLarge,
"JSON request body exceeds this route's body limit",
)
}
pub(super) struct OptionalAppJson<T>(pub(super) Option<T>);
const MAX_OPTIONAL_JSON_BODY_BYTES: usize = 1024 * 1024;
impl<T> FromRequest<AppState> for OptionalAppJson<T>
where
T: serde::de::DeserializeOwned,
{
type Rejection = ApiResponseError;
async fn from_request(
req: axum::extract::Request,
state: &AppState,
) -> Result<Self, Self::Rejection> {
authorize(&state.config, req.headers())?;
let body = axum::body::to_bytes(req.into_body(), MAX_OPTIONAL_JSON_BODY_BYTES)
.await
.map_err(|error| {
ApiResponseError::new(
StatusCode::BAD_REQUEST,
ErrorCode::InvalidRequest,
&format!("request body unreadable: {error}"),
)
})?;
if body.is_empty() {
return Ok(OptionalAppJson(None));
}
let value = serde_json::from_slice(&body).map_err(|error| {
ApiResponseError::new(
StatusCode::BAD_REQUEST,
ErrorCode::InvalidRequest,
&format!("request body is not valid JSON for this operation: {error}"),
)
})?;
Ok(OptionalAppJson(Some(value)))
}
}