This commit is contained in:
David Pedersen
2022-03-13 22:56:19 +01:00
parent 630287f76a
commit 2e7169e716
9 changed files with 531 additions and 114 deletions
+3
View File
@@ -150,6 +150,9 @@ pub trait TypedPath: std::fmt::Display {
/// The path with optional captures such as `/users/:id`. /// The path with optional captures such as `/users/:id`.
const PATH: &'static str; const PATH: &'static str;
/// The parameter types this path requires.
type Parameters;
/// Convert the path into a `Uri`. /// Convert the path into a `Uri`.
/// ///
/// # Panics /// # Panics
+18 -3
View File
@@ -21,9 +21,9 @@ pub(crate) fn expand(item_struct: ItemStruct) -> syn::Result<TokenStream> {
let Attrs { path } = parse_attrs(attrs)?; let Attrs { path } = parse_attrs(attrs)?;
match fields { match fields {
syn::Fields::Named(_) => { syn::Fields::Named(fields) => {
let segments = parse_path(&path)?; let segments = parse_path(&path)?;
Ok(expand_named_fields(ident, path, &segments)) Ok(expand_named_fields(ident, fields, path, &segments))
} }
syn::Fields::Unnamed(fields) => { syn::Fields::Unnamed(fields) => {
let segments = parse_path(&path)?; let segments = parse_path(&path)?;
@@ -63,14 +63,23 @@ fn parse_attrs(attrs: &[syn::Attribute]) -> syn::Result<Attrs> {
}) })
} }
fn expand_named_fields(ident: &syn::Ident, path: LitStr, segments: &[Segment]) -> TokenStream { fn expand_named_fields(
ident: &syn::Ident,
fields: &syn::FieldsNamed,
path: LitStr,
segments: &[Segment],
) -> TokenStream {
let format_str = format_str_from_path(segments); let format_str = format_str_from_path(segments);
let captures = captures_from_path(segments); let captures = captures_from_path(segments);
let params = fields.named.iter().map(|field| &field.ty);
let typed_path_impl = quote_spanned! {path.span()=> let typed_path_impl = quote_spanned! {path.span()=>
#[automatically_derived] #[automatically_derived]
impl ::axum_extra::routing::TypedPath for #ident { impl ::axum_extra::routing::TypedPath for #ident {
const PATH: &'static str = #path; const PATH: &'static str = #path;
type Parameters = (#(#params,)*);
} }
}; };
@@ -156,10 +165,14 @@ fn expand_unnamed_fields(
let format_str = format_str_from_path(segments); let format_str = format_str_from_path(segments);
let captures = captures_from_path(segments); let captures = captures_from_path(segments);
let params = fields.unnamed.iter().map(|field| &field.ty);
let typed_path_impl = quote_spanned! {path.span()=> let typed_path_impl = quote_spanned! {path.span()=>
#[automatically_derived] #[automatically_derived]
impl ::axum_extra::routing::TypedPath for #ident { impl ::axum_extra::routing::TypedPath for #ident {
const PATH: &'static str = #path; const PATH: &'static str = #path;
type Parameters = (#(#params,)*);
} }
}; };
@@ -224,6 +237,8 @@ fn expand_unit_fields(ident: &syn::Ident, path: LitStr) -> syn::Result<TokenStre
#[automatically_derived] #[automatically_derived]
impl ::axum_extra::routing::TypedPath for #ident { impl ::axum_extra::routing::TypedPath for #ident {
const PATH: &'static str = #path; const PATH: &'static str = #path;
type Parameters = ();
} }
}; };
+7
View File
@@ -7,3 +7,10 @@ edition = "2021"
axum = { path = "../axum", version = "0.4" } axum = { path = "../axum", version = "0.4" }
okapi = "0.7.0-rc.1" okapi = "0.7.0-rc.1"
schemars = "0.8.8" schemars = "0.8.8"
mime = "0.3.16"
serde_json = "1.0"
[dev-dependencies]
assert-json-diff = "2.0"
serde = { version = "1.0", features = ["derive"] }
+91
View File
@@ -0,0 +1,91 @@
use axum::{
body::HttpBody,
extract::FromRequest,
handler::Handler,
http::{Request, StatusCode},
response::{IntoResponse, Response},
routing::{self, MethodRouter},
Json, Router,
};
use okapi::openapi3::{
self, Components, Info, MediaType, OpenApi, Operation, Parameter, RefOr, RequestBody,
};
use schemars::{
schema::{RootSchema, Schema},
JsonSchema,
};
use std::{
collections::BTreeMap, convert::Infallible, future::Future, marker::PhantomData, sync::Arc,
};
pub trait DescribeRequest {
fn describe(operation: &mut Operation, components: &mut Components);
}
impl DescribeRequest for () {
fn describe(_: &mut Operation, _: &mut Components) {}
}
macro_rules! impl_tuples {
( $($ty:ident),* $(,)? ) => {
impl<$($ty,)*> DescribeRequest for ($($ty,)*)
where
$($ty: DescribeRequest,)*
{
fn describe(operation: &mut Operation, components: &mut Components) {
$( $ty::describe(operation, components); )*
}
}
};
}
all_the_tuples!(impl_tuples);
impl<T> DescribeRequest for Json<T>
where
T: JsonSchema,
{
fn describe(operation: &mut Operation, components: &mut Components) {
let RootSchema {
mut schema,
definitions,
meta_schema: _,
} = schemars::schema_for!(T);
components.schemas.extend(
definitions
.into_iter()
.filter_map(|(k, schema)| match schema {
Schema::Bool(_) => None,
Schema::Object(obj) => Some((k, obj)),
}),
);
schema.object().properties = std::mem::take(&mut schema.object().properties)
.into_iter()
.map(|(key, schema)| match schema {
Schema::Bool(_) => (key, schema),
Schema::Object(mut obj) => {
if let Some(reference) = &mut obj.reference {
*reference = reference.replace("/definitions/", "/components/schemas/");
}
(key, Schema::Object(obj))
}
})
.collect();
let request_body = RequestBody {
content: schemars::Map::from_iter([(
mime::APPLICATION_JSON.to_string(),
MediaType {
schema: Some(schema),
..Default::default()
},
)]),
required: true,
..Default::default()
};
operation.request_body = Some(RefOr::Object(request_body));
}
}
+74
View File
@@ -0,0 +1,74 @@
use axum::{
body::HttpBody,
extract::FromRequest,
handler::Handler,
http::{Request, StatusCode},
response::{IntoResponse, Response},
routing::{self, MethodRouter},
Json, Router,
};
use okapi::openapi3::{
self, Components, Info, MediaType, OpenApi, Operation, Parameter, RefOr, RequestBody,
};
use schemars::{
schema::{RootSchema, Schema},
JsonSchema,
};
use std::{
collections::BTreeMap, convert::Infallible, future::Future, marker::PhantomData, sync::Arc,
};
pub trait DescribeResponse {
fn describe(operation: &mut Operation, components: &mut Components);
}
impl DescribeResponse for () {
fn describe(operation: &mut Operation, components: &mut Components) {
Ok::describe(operation, components)
}
}
macro_rules! impl_tuples {
( $($ty:ident),* $(,)? ) => {
impl<$($ty,)*> DescribeResponse for ($($ty,)*)
where
$($ty: DescribeResponse,)*
{
fn describe(operation: &mut Operation, components: &mut Components) {
$( $ty::describe(operation, components); )*
}
}
};
}
all_the_tuples!(impl_tuples);
macro_rules! status {
(
$name:ident, $variant:ident
) => {
#[derive(Copy, Clone)]
pub struct $name;
impl IntoResponse for $name {
fn into_response(self) -> Response {
StatusCode::$variant.into_response()
}
}
impl DescribeResponse for $name {
fn describe(operation: &mut Operation, _: &mut Components) {
operation.responses.responses.insert(
StatusCode::$variant.as_u16().to_string(),
RefOr::Object(openapi3::Response {
description: "Successful response".to_owned(),
..Default::default()
}),
);
}
}
};
}
status!(Ok, OK);
status!(Created, CREATED);
+291 -103
View File
@@ -1,16 +1,57 @@
#![allow(missing_debug_implementations)] #![allow(missing_debug_implementations, dead_code, unused_imports)]
#![deny(unreachable_pub)]
use axum::{ use axum::{
async_trait,
body::HttpBody, body::HttpBody,
extract::FromRequest,
handler::Handler, handler::Handler,
http::Request, http::{Request, StatusCode},
response::{IntoResponse, Response}, response::{IntoResponse, Response},
routing::{get, post}, routing::{self, MethodRouter},
Router, Json, Router,
};
use okapi::openapi3::{
self, Components, Info, MediaType, OpenApi, Operation, Parameter, RefOr, RequestBody,
};
use schemars::{
schema::{RootSchema, Schema},
JsonSchema,
};
use std::{
collections::BTreeMap, convert::Infallible, future::Future, marker::PhantomData, sync::Arc,
};
#[macro_use]
mod macros {
macro_rules! all_the_tuples {
($name:ident) => {
$name!(T1);
$name!(T1, T2);
$name!(T1, T2, T3);
$name!(T1, T2, T3, T4);
$name!(T1, T2, T3, T4, T5);
$name!(T1, T2, T3, T4, T5, T6);
$name!(T1, T2, T3, T4, T5, T6, T7);
$name!(T1, T2, T3, T4, T5, T6, T7, T8);
$name!(T1, T2, T3, T4, T5, T6, T7, T8, T9);
$name!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10);
$name!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11);
$name!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12);
$name!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13);
$name!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13, T14);
$name!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13, T14, T15);
$name!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13, T14, T15, T16);
};
}
}
mod describe_request;
mod describe_response;
pub use self::{
describe_request::DescribeRequest,
describe_response::{Created, DescribeResponse, Ok},
}; };
use okapi::openapi3::{Info, OpenApi, Operation};
use std::{future::Future, marker::PhantomData, sync::Arc};
pub struct OpenApiRouter<B> { pub struct OpenApiRouter<B> {
router: Router<B>, router: Router<B>,
@@ -26,129 +67,276 @@ where
router: Default::default(), router: Default::default(),
schema: OpenApi { schema: OpenApi {
info, info,
openapi: "3.0.0".to_owned(),
components: Some(Components::default()),
..Default::default() ..Default::default()
}, },
} }
} }
pub fn get<H, T>(mut self, path: &str, handler: H) -> Self pub fn into_parts(self) -> (Router<B>, OpenApi) {
where (self.router, self.schema)
H: OpenApiHandler<T, B>,
T: 'static,
{
let mut operation = Operation::default();
handler.clone().to_operation(&mut operation);
self.schema.paths.entry(path.to_owned()).or_default().get = Some(operation);
self.router = self.router.route(path, get(handler));
self
} }
pub fn post<H, T>(mut self, path: &str, handler: H) -> Self pub fn route(mut self, path: &str, handler: OpenApiHandler<B>) -> Self {
where let OpenApiHandler {
H: OpenApiHandler<T, B>, svc,
T: 'static, method,
{ operation,
let mut operation = Operation::default(); components,
handler.clone().to_operation(&mut operation); } = handler;
self.schema.paths.entry(path.to_owned()).or_default().post = Some(operation);
self.router = self.router.route(path, post(handler)); self.router = self.router.route(path, svc);
extend_components(self.schema.components.as_mut().unwrap(), components);
let path = path
.split('/')
.map(|segment| {
if let Some(param) = segment.strip_prefix(':') {
format!("{{{}}}", param)
} else {
// TODO(david): wildcards
segment.to_owned()
}
})
.collect::<Vec<_>>()
.join("/");
let path_item = self.schema.paths.entry(path).or_default();
match method {
Method::Get => path_item.get = Some(operation),
Method::Post => path_item.post = Some(operation),
}
self self
} }
} }
pub trait OpenApiHandler<T, B>: Handler<T, B> { fn extend_components(current: &mut Components, new: Components) {
fn to_operation(self, operation: &mut Operation); let Components {
schemas,
responses,
parameters,
examples,
request_bodies,
headers,
security_schemes,
links,
callbacks,
extensions,
} = current;
fn map_operation<F>(self, f: F) -> MapOperation<Self, T, B, F> schemas.extend(new.schemas);
responses.extend(new.responses);
parameters.extend(new.parameters);
examples.extend(new.examples);
request_bodies.extend(new.request_bodies);
headers.extend(new.headers);
security_schemes.extend(new.security_schemes);
links.extend(new.links);
callbacks.extend(new.callbacks);
extensions.extend(new.extensions);
}
macro_rules! method {
($fn_name:ident, $method:ident) => {
pub fn $fn_name<H, T, B>(handler: H) -> OpenApiHandler<B>
where
H: Handler<T, B> + HandlerResponse<T>,
T: DescribeRequest + 'static,
H::Response: DescribeResponse,
B: Send + 'static,
{
let mut operation = Operation::default();
let mut components = Components::default();
T::describe(&mut operation, &mut components);
H::Response::describe(&mut operation, &mut components);
OpenApiHandler {
svc: routing::$fn_name(handler),
method: Method::$method,
operation,
components,
}
}
};
}
method!(get, Get);
method!(post, Post);
// TODO(david): the remaining methods
enum Method {
Get,
Post,
}
pub struct OpenApiHandler<B> {
svc: MethodRouter<B, Infallible>,
method: Method,
operation: Operation,
components: Components,
}
impl<B> OpenApiHandler<B> {
pub fn operation_id<S>(self, id: S) -> Self
where where
F: FnOnce(&mut Operation), S: Into<String>,
{ {
MapOperation { self.map_operation(|op, _| {
handler: self, op.operation_id = Some(id.into());
f, })
_marker: PhantomData, }
pub fn summary<S>(self, summary: S) -> Self
where
S: Into<String>,
{
self.map_operation(|op, _| {
op.summary = Some(summary.into());
})
}
pub fn description<S>(self, description: S) -> Self
where
S: Into<String>,
{
self.map_operation(|op, _| {
op.description = Some(description.into());
})
}
pub fn map_operation<F>(mut self, f: F) -> Self
where
F: FnOnce(&mut Operation, &mut Components),
{
f(&mut self.operation, &mut self.components);
self
}
}
pub trait HandlerResponse<T> {
type Response;
}
impl<F, Fut> HandlerResponse<()> for F
where
F: FnOnce() -> Fut,
Fut: Future,
{
type Response = Fut::Output;
}
macro_rules! impl_tuples {
( $($ty:ident),* $(,)? ) => {
impl<F, Fut, $($ty,)*> HandlerResponse<($($ty,)*)> for F
where
F: FnOnce($($ty,)*) -> Fut,
Fut: Future,
{
type Response = Fut::Output;
} }
} };
} }
pub struct MapOperation<H, T, B, F> { all_the_tuples!(impl_tuples);
handler: H,
f: F,
_marker: PhantomData<(T, B)>,
}
impl<H, T, B, F> Clone for MapOperation<H, T, B, F>
where
H: Clone,
F: Clone,
{
fn clone(&self) -> Self {
Self {
handler: self.handler.clone(),
f: self.f.clone(),
_marker: PhantomData,
}
}
}
impl<H, T, B, F> OpenApiHandler<T, B> for MapOperation<H, T, B, F>
where
H: OpenApiHandler<T, B> + Handler<T, B>,
F: FnOnce(&mut Operation) + Clone + Send + 'static,
T: Send + 'static,
B: Send + 'static,
{
fn to_operation(self, operation: &mut Operation) {
(self.f)(operation);
}
}
impl<H, T, B, F> Handler<T, B> for MapOperation<H, T, B, F>
where
H: Handler<T, B>,
F: Clone + Send + 'static,
T: Send + 'static,
B: Send + 'static,
{
type Future = H::Future;
fn call(self, req: Request<B>) -> Self::Future {
self.handler.call(req)
}
}
impl<F, Fut, Res, B> OpenApiHandler<(), B> for F
where
F: FnOnce() -> Fut + Clone + Send + 'static,
Fut: Future<Output = Res> + Send,
Res: IntoResponse,
B: Send + 'static,
{
fn to_operation(self, operation: &mut Operation) {
todo!()
}
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use assert_json_diff::assert_json_eq;
use axum::body::Body; use axum::body::Body;
use serde::Deserialize;
use serde_json::json;
#[test] #[test]
fn test_something() { fn test_something() {
#[derive(Deserialize, JsonSchema)]
struct UsersCreate {
account: Account,
}
#[derive(Deserialize, JsonSchema)]
struct Account {
username: String,
}
async fn users_show() {} async fn users_show() {}
async fn users_create() {} async fn users_create(Json(_): Json<UsersCreate>) -> Created {
Created
}
let _router: OpenApiRouter<Body> = OpenApiRouter::new(Info::default()) let (router, schema) = OpenApiRouter::<Body>::new(Info::default())
.post( .route("/users/:id", get(users_show).operation_id("users_show"))
"/users", .route("/users", post(users_create).operation_id("users_create"))
users_create .into_parts();
.map_operation(|_| {})
.map_operation(|mut operation| { assert_json_eq!(
operation.summary = Some("Create a new user".to_owned()); schema,
}), json!({
) "openapi": "3.0.0",
.get("/users/:id", users_show.map_operation(|_operation| {})); "info": {
"title": "",
"version": ""
},
"paths": {
"/users": {
"post": {
"operationId": "users_create",
"requestBody": {
"content": {
"application/json": {
"schema": {
"title": "UsersCreate",
"type": "object",
"required": [
"account",
],
"properties": {
"account": {
"$ref": "#/components/schemas/Account"
}
}
}
}
},
"required": true
},
"responses": {
"201": {
"description": "Successful response",
}
}
}
},
"/users/{id}": {
"get": {
"operationId": "users_show",
"responses": {
"200": {
"description": "Successful response",
}
}
}
}
},
"components": {
"schemas": {
"Account": {
"type": "object",
"required": [
"username"
],
"properties": {
"username": {
"type": "string"
}
}
}
}
}
}),
);
} }
} }
+38 -1
View File
@@ -2,7 +2,11 @@
use crate::response::Response; use crate::response::Response;
use futures_util::future::Map; use futures_util::future::Map;
use std::convert::Infallible; use http::Request;
use pin_project_lite::pin_project;
use std::{convert::Infallible, future::Future, pin::Pin, task::Context};
use tower::util::Oneshot;
use tower_service::Service;
opaque_future! { opaque_future! {
/// The response future for [`IntoService`](super::IntoService). /// The response future for [`IntoService`](super::IntoService).
@@ -12,3 +16,36 @@ opaque_future! {
fn(Response) -> Result<Response, Infallible>, fn(Response) -> Result<Response, Infallible>,
>; >;
} }
pin_project! {
/// The response future for [`Layered`](super::Layered).
pub struct LayeredFuture<S, ReqBody>
where
S: Service<Request<ReqBody>>,
{
#[pin]
inner: Map<Oneshot<S, Request<ReqBody>>, fn(Result<S::Response, S::Error>) -> Response>,
}
}
impl<S, ReqBody> LayeredFuture<S, ReqBody>
where
S: Service<Request<ReqBody>>,
{
pub(super) fn new(
inner: Map<Oneshot<S, Request<ReqBody>>, fn(Result<S::Response, S::Error>) -> Response>,
) -> Self {
Self { inner }
}
}
impl<S, ReqBody> Future for LayeredFuture<S, ReqBody>
where
S: Service<Request<ReqBody>>,
{
type Output = Response;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> std::task::Poll<Self::Output> {
self.project().inner.poll(cx)
}
}
+9 -5
View File
@@ -99,6 +99,7 @@ pub use self::into_service::IntoService;
/// ///
/// See the [module docs](crate::handler) for more details. /// See the [module docs](crate::handler) for more details.
pub trait Handler<T, B = Body>: Clone + Send + Sized + 'static { pub trait Handler<T, B = Body>: Clone + Send + Sized + 'static {
/// The type of future calling this handler returns.
type Future: Future<Output = Response> + Send + 'static; type Future: Future<Output = Response> + Send + 'static;
/// Call the handler with the given request. /// Call the handler with the given request.
@@ -334,15 +335,18 @@ where
ResBody: HttpBody<Data = Bytes> + Send + 'static, ResBody: HttpBody<Data = Bytes> + Send + 'static,
ResBody::Error: Into<BoxError>, ResBody::Error: Into<BoxError>,
{ {
type Future = Pin<Box<dyn Future<Output = Response> + Send + 'static>>; type Future = future::LayeredFuture<S, ReqBody>;
fn call(self, req: Request<ReqBody>) -> Self::Future { fn call(self, req: Request<ReqBody>) -> Self::Future {
Box::pin(async move { use futures_util::future::{FutureExt, Map};
match self.svc.oneshot(req).await {
let future: Map<_, fn(Result<S::Response, S::Error>) -> _> =
self.svc.oneshot(req).map(|result| match result {
Ok(res) => res.map(boxed), Ok(res) => res.map(boxed),
Err(res) => res.into_response(), Err(res) => res.into_response(),
} });
})
future::LayeredFuture::new(future)
} }
} }
-2
View File
@@ -1,5 +1,3 @@
//! Internal macros
macro_rules! opaque_future { macro_rules! opaque_future {
($(#[$m:meta])* pub type $name:ident = $actual:ty;) => { ($(#[$m:meta])* pub type $name:ident = $actual:ty;) => {
opaque_future! { opaque_future! {