mirror of
https://github.com/tokio-rs/axum.git
synced 2026-09-07 00:00:12 +02:00
axum: use Vec for PathRouter (#3509)
This commit is contained in:
@@ -58,7 +58,7 @@ macro_rules! panic_on_err {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
|
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
|
||||||
pub(crate) struct RouteId(u32);
|
pub(crate) struct RouteId(usize);
|
||||||
|
|
||||||
/// The router type for composing handlers and services.
|
/// The router type for composing handlers and services.
|
||||||
///
|
///
|
||||||
|
|||||||
@@ -14,9 +14,8 @@ use super::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
pub(super) struct PathRouter<S> {
|
pub(super) struct PathRouter<S> {
|
||||||
routes: HashMap<RouteId, Endpoint<S>>,
|
routes: Vec<Endpoint<S>>,
|
||||||
node: Arc<Node>,
|
node: Arc<Node>,
|
||||||
prev_route_id: RouteId,
|
|
||||||
v7_checks: bool,
|
v7_checks: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -71,11 +70,11 @@ where
|
|||||||
) -> Result<(), Cow<'static, str>> {
|
) -> Result<(), Cow<'static, str>> {
|
||||||
validate_path(self.v7_checks, path)?;
|
validate_path(self.v7_checks, path)?;
|
||||||
|
|
||||||
let endpoint = if let Some((route_id, Endpoint::MethodRouter(prev_method_router))) = self
|
if let Some((route_id, Endpoint::MethodRouter(prev_method_router))) = self
|
||||||
.node
|
.node
|
||||||
.path_to_route_id
|
.path_to_route_id
|
||||||
.get(path)
|
.get(path)
|
||||||
.and_then(|route_id| self.routes.get(route_id).map(|svc| (*route_id, svc)))
|
.and_then(|route_id| self.routes.get(route_id.0).map(|svc| (*route_id, svc)))
|
||||||
{
|
{
|
||||||
// if we're adding a new `MethodRouter` to a route that already has one just
|
// if we're adding a new `MethodRouter` to a route that already has one just
|
||||||
// merge them. This makes `.route("/", get(_)).route("/", post(_))` work
|
// merge them. This makes `.route("/", get(_)).route("/", post(_))` work
|
||||||
@@ -84,15 +83,11 @@ where
|
|||||||
.clone()
|
.clone()
|
||||||
.merge_for_path(Some(path), method_router)?,
|
.merge_for_path(Some(path), method_router)?,
|
||||||
);
|
);
|
||||||
self.routes.insert(route_id, service);
|
self.routes[route_id.0] = service;
|
||||||
return Ok(());
|
|
||||||
} else {
|
} else {
|
||||||
Endpoint::MethodRouter(method_router)
|
let endpoint = Endpoint::MethodRouter(method_router);
|
||||||
};
|
self.new_route(path, endpoint)?;
|
||||||
|
}
|
||||||
let id = self.next_route_id();
|
|
||||||
self.set_node(path, id)?;
|
|
||||||
self.routes.insert(id, endpoint);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
@@ -102,7 +97,7 @@ where
|
|||||||
H: Handler<T, S>,
|
H: Handler<T, S>,
|
||||||
T: 'static,
|
T: 'static,
|
||||||
{
|
{
|
||||||
for (_, endpoint) in self.routes.iter_mut() {
|
for endpoint in self.routes.iter_mut() {
|
||||||
if let Endpoint::MethodRouter(rt) = endpoint {
|
if let Endpoint::MethodRouter(rt) = endpoint {
|
||||||
*rt = rt.clone().default_fallback(handler.clone());
|
*rt = rt.clone().default_fallback(handler.clone());
|
||||||
}
|
}
|
||||||
@@ -129,9 +124,7 @@ where
|
|||||||
) -> Result<(), Cow<'static, str>> {
|
) -> Result<(), Cow<'static, str>> {
|
||||||
validate_path(self.v7_checks, path)?;
|
validate_path(self.v7_checks, path)?;
|
||||||
|
|
||||||
let id = self.next_route_id();
|
self.new_route(path, endpoint)?;
|
||||||
self.set_node(path, id)?;
|
|
||||||
self.routes.insert(id, endpoint);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
@@ -143,21 +136,28 @@ where
|
|||||||
.map_err(|err| format!("Invalid route {path:?}: {err}"))
|
.map_err(|err| format!("Invalid route {path:?}: {err}"))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn new_route(&mut self, path: &str, endpoint: Endpoint<S>) -> Result<(), String> {
|
||||||
|
let id = RouteId(self.routes.len());
|
||||||
|
self.set_node(path, id)?;
|
||||||
|
self.routes.push(endpoint);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
pub(super) fn merge(&mut self, other: Self) -> Result<(), Cow<'static, str>> {
|
pub(super) fn merge(&mut self, other: Self) -> Result<(), Cow<'static, str>> {
|
||||||
let Self {
|
let Self {
|
||||||
routes,
|
routes,
|
||||||
node,
|
node,
|
||||||
prev_route_id: _,
|
|
||||||
v7_checks,
|
v7_checks,
|
||||||
} = other;
|
} = other;
|
||||||
|
|
||||||
// If either of the two did not allow paths starting with `:` or `*`, do not allow them for the merged router either.
|
// If either of the two did not allow paths starting with `:` or `*`, do not allow them for the merged router either.
|
||||||
self.v7_checks |= v7_checks;
|
self.v7_checks |= v7_checks;
|
||||||
|
|
||||||
for (id, route) in routes {
|
for (id, route) in routes.into_iter().enumerate() {
|
||||||
|
let route_id = RouteId(id);
|
||||||
let path = node
|
let path = node
|
||||||
.route_id_to_path
|
.route_id_to_path
|
||||||
.get(&id)
|
.get(&route_id)
|
||||||
.expect("no path for route id. This is a bug in axum. Please file an issue");
|
.expect("no path for route id. This is a bug in axum. Please file an issue");
|
||||||
|
|
||||||
match route {
|
match route {
|
||||||
@@ -179,15 +179,15 @@ where
|
|||||||
let Self {
|
let Self {
|
||||||
routes,
|
routes,
|
||||||
node,
|
node,
|
||||||
prev_route_id: _,
|
|
||||||
// Ignore the configuration of the nested router
|
// Ignore the configuration of the nested router
|
||||||
v7_checks: _,
|
v7_checks: _,
|
||||||
} = router;
|
} = router;
|
||||||
|
|
||||||
for (id, endpoint) in routes {
|
for (id, endpoint) in routes.into_iter().enumerate() {
|
||||||
|
let route_id = RouteId(id);
|
||||||
let inner_path = node
|
let inner_path = node
|
||||||
.route_id_to_path
|
.route_id_to_path
|
||||||
.get(&id)
|
.get(&route_id)
|
||||||
.expect("no path for route id. This is a bug in axum. Please file an issue");
|
.expect("no path for route id. This is a bug in axum. Please file an issue");
|
||||||
|
|
||||||
let path = path_for_nested_route(prefix, inner_path);
|
let path = path_for_nested_route(prefix, inner_path);
|
||||||
@@ -259,16 +259,12 @@ where
|
|||||||
let routes = self
|
let routes = self
|
||||||
.routes
|
.routes
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.map(|(id, endpoint)| {
|
.map(|endpoint| endpoint.layer(layer.clone()))
|
||||||
let route = endpoint.layer(layer.clone());
|
|
||||||
(id, route)
|
|
||||||
})
|
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
Self {
|
Self {
|
||||||
routes,
|
routes,
|
||||||
node: self.node,
|
node: self.node,
|
||||||
prev_route_id: self.prev_route_id,
|
|
||||||
v7_checks: self.v7_checks,
|
v7_checks: self.v7_checks,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -292,16 +288,12 @@ where
|
|||||||
let routes = self
|
let routes = self
|
||||||
.routes
|
.routes
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.map(|(id, endpoint)| {
|
.map(|endpoint| endpoint.layer(layer.clone()))
|
||||||
let route = endpoint.layer(layer.clone());
|
|
||||||
(id, route)
|
|
||||||
})
|
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
Self {
|
Self {
|
||||||
routes,
|
routes,
|
||||||
node: self.node,
|
node: self.node,
|
||||||
prev_route_id: self.prev_route_id,
|
|
||||||
v7_checks: self.v7_checks,
|
v7_checks: self.v7_checks,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -314,21 +306,17 @@ where
|
|||||||
let routes = self
|
let routes = self
|
||||||
.routes
|
.routes
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.map(|(id, endpoint)| {
|
.map(|endpoint| match endpoint {
|
||||||
let endpoint: Endpoint<S2> = match endpoint {
|
Endpoint::MethodRouter(method_router) => {
|
||||||
Endpoint::MethodRouter(method_router) => {
|
Endpoint::MethodRouter(method_router.with_state(state.clone()))
|
||||||
Endpoint::MethodRouter(method_router.with_state(state.clone()))
|
}
|
||||||
}
|
Endpoint::Route(route) => Endpoint::Route(route),
|
||||||
Endpoint::Route(route) => Endpoint::Route(route),
|
|
||||||
};
|
|
||||||
(id, endpoint)
|
|
||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
PathRouter {
|
PathRouter {
|
||||||
routes,
|
routes,
|
||||||
node: self.node,
|
node: self.node,
|
||||||
prev_route_id: self.prev_route_id,
|
|
||||||
v7_checks: self.v7_checks,
|
v7_checks: self.v7_checks,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -366,7 +354,7 @@ where
|
|||||||
|
|
||||||
let endpoint = self
|
let endpoint = self
|
||||||
.routes
|
.routes
|
||||||
.get(&id)
|
.get(id.0)
|
||||||
.expect("no route for id. This is a bug in axum. Please file an issue");
|
.expect("no route for id. This is a bug in axum. Please file an issue");
|
||||||
|
|
||||||
let req = Request::from_parts(parts, body);
|
let req = Request::from_parts(parts, body);
|
||||||
@@ -382,16 +370,6 @@ where
|
|||||||
Err(MatchError::NotFound) => Err((Request::from_parts(parts, body), state)),
|
Err(MatchError::NotFound) => Err((Request::from_parts(parts, body), state)),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn next_route_id(&mut self) -> RouteId {
|
|
||||||
let next_id = self
|
|
||||||
.prev_route_id
|
|
||||||
.0
|
|
||||||
.checked_add(1)
|
|
||||||
.expect("Over `u32::MAX` routes created. If you need this, please file an issue.");
|
|
||||||
self.prev_route_id = RouteId(next_id);
|
|
||||||
self.prev_route_id
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<S> Default for PathRouter<S> {
|
impl<S> Default for PathRouter<S> {
|
||||||
@@ -399,7 +377,6 @@ impl<S> Default for PathRouter<S> {
|
|||||||
Self {
|
Self {
|
||||||
routes: Default::default(),
|
routes: Default::default(),
|
||||||
node: Default::default(),
|
node: Default::default(),
|
||||||
prev_route_id: RouteId(0),
|
|
||||||
v7_checks: true,
|
v7_checks: true,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -419,7 +396,6 @@ impl<S> Clone for PathRouter<S> {
|
|||||||
Self {
|
Self {
|
||||||
routes: self.routes.clone(),
|
routes: self.routes.clone(),
|
||||||
node: self.node.clone(),
|
node: self.node.clone(),
|
||||||
prev_route_id: self.prev_route_id,
|
|
||||||
v7_checks: self.v7_checks,
|
v7_checks: self.v7_checks,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user