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/spec/vtc/install/claim/0.1").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
}
#[utoipa::path(get, path = "/v1/admin/config", responses((status = 200)))]
async fn cfg_show() -> &'static str {
"show"
}
#[utoipa::path(patch, path = "/v1/admin/config", responses((status = 200)))]
async fn cfg_patch() -> &'static str {
"patch"
}
const SHOW: &str = "https://trusttasks.org/spec/config/show/0.1";
const PATCH: &str = "https://trusttasks.org/spec/config/patch/0.1";
async fn split_app() -> Router {
OpenApiRouter::new()
.routes(task_routes(
routes!(cfg_show),
TrustTask::new(SHOW).unwrap(),
))
.routes(task_routes(
routes!(cfg_patch),
TrustTask::new(PATCH).unwrap(),
))
.split_for_parts()
.0
}
async fn call(method: &str, task: &str) -> StatusCode {
split_app()
.await
.oneshot(
Request::builder()
.method(method)
.uri("/v1/admin/config")
.header(HEADER_NAME, task)
.body(Body::empty())
.unwrap(),
)
.await
.unwrap()
.status()
}
#[test]
fn per_method_split_keeps_both_operations_in_the_spec() {
let (_router, api): (Router, _) = OpenApiRouter::new()
.routes(task_routes(
routes!(cfg_show),
TrustTask::new(SHOW).unwrap(),
))
.routes(task_routes(
routes!(cfg_patch),
TrustTask::new(PATCH).unwrap(),
))
.split_for_parts();
let item = api
.paths
.paths
.get("/v1/admin/config")
.expect("path present");
assert!(item.get.is_some(), "GET operation dropped from the spec");
assert!(
item.patch.is_some(),
"PATCH operation dropped from the spec"
);
}
#[tokio::test]
async fn per_method_tasks_on_one_path_are_enforced_independently() {
assert_eq!(call("GET", SHOW).await, StatusCode::OK);
assert_eq!(call("PATCH", PATCH).await, StatusCode::OK);
assert_eq!(call("GET", PATCH).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
assert_eq!(
call("PATCH", SHOW).await,
StatusCode::UNSUPPORTED_MEDIA_TYPE
);
}
#[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/spec/vtc/install/claim/0.1",
)
.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/spec/vtc/auth/login/0.1",
)
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNSUPPORTED_MEDIA_TYPE);
}
}