use std::sync::Arc;
use axum::routing::MethodRouter;
use utoipa_axum::router::{UtoipaMethodRouter, UtoipaMethodRouterExt};
use super::TrustTask;
pub fn task_routes<S>(routes: UtoipaMethodRouter<S>, task: TrustTask) -> UtoipaMethodRouter<S>
where
S: Clone + Send + Sync + 'static,
{
let task = Arc::new(task);
routes.layer(axum::middleware::from_fn(move |request, next| {
let task = task.clone();
async move { super::extractor::validate_header(&task, request, next).await }
}))
}
pub fn task_layer<S>(method_router: MethodRouter<S>, task: TrustTask) -> MethodRouter<S>
where
S: Clone + Send + Sync + 'static,
{
let task = Arc::new(task);
method_router.layer(axum::middleware::from_fn(move |request, next| {
let task = task.clone();
async move { super::extractor::validate_header(&task, request, next).await }
}))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::trust_task::HEADER_NAME;
use axum::Router;
use axum::body::Body;
use axum::http::{Request, StatusCode};
use tower::ServiceExt;
use utoipa_axum::router::OpenApiRouter;
use utoipa_axum::routes;
#[utoipa::path(get, path = "/v1/install/claim", responses((status = 200)))]
async fn ok() -> &'static str {
"ok"
}
fn task() -> TrustTask {
TrustTask::new("https://trusttasks.org/openvtc/vtc/install/claim/1.0").unwrap()
}
#[test]
fn task_routes_keeps_the_operation_in_the_spec() {
let (_router, api): (Router, _) = OpenApiRouter::new()
.routes(task_routes(routes!(ok), task()))
.split_for_parts();
assert!(api.paths.paths.contains_key("/v1/install/claim"));
}
async fn app() -> Router {
OpenApiRouter::new()
.routes(task_routes(routes!(ok), task()))
.split_for_parts()
.0
}
#[tokio::test]
async fn task_routes_enforces_the_header() {
let resp = app()
.await
.oneshot(
Request::builder()
.uri("/v1/install/claim")
.header(
HEADER_NAME,
"https://trusttasks.org/openvtc/vtc/install/claim/1.0",
)
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let resp = app()
.await
.oneshot(
Request::builder()
.uri("/v1/install/claim")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let resp = app()
.await
.oneshot(
Request::builder()
.uri("/v1/install/claim")
.header(
HEADER_NAME,
"https://trusttasks.org/openvtc/vtc/auth/login/1.0",
)
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNSUPPORTED_MEDIA_TYPE);
}
}