diff --git a/Cargo.toml b/Cargo.toml index 81606090..26d50a56 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -16,6 +16,7 @@ clokwerk = "^0.1" derive_more = "^0.13" fern = { version = "^0.5", features = ["syslog-4"] } futures = "0.3.4" +futures-util = "0.3.4" hex = "^0.3" hyper = "^0.13" log = "^0.4" diff --git a/src/daemon/endpoints.rs b/src/daemon/endpoints.rs index caebd4fe..23118691 100644 --- a/src/daemon/endpoints.rs +++ b/src/daemon/endpoints.rs @@ -39,6 +39,13 @@ fn render_json_res(res: Result) -> RoutingResult { } } +fn render_json_res_res(res: Result, Error>) -> RoutingResult { + match res { + Ok(res) => render_json_res(res), + Err(e) => render_error(e), + } +} + /// A clean 404 result for the API (no content, not for humans) fn render_unknown_resource() -> RoutingResult { Ok(HttpResponse::error(Error::ApiUnknownResource)) @@ -54,12 +61,12 @@ fn render_unknown_method() -> RoutingResult { } /// A clean 404 response -pub fn render_not_found(_req: Request) -> RoutingResult { +pub async fn render_not_found(_req: Request) -> RoutingResult { Ok(HttpResponse::not_found()) } /// Returns the server health. -pub fn health(req: Request) -> RoutingResult { +pub async fn health(req: Request) -> RoutingResult { if req.is_get() && req.path().segment() == "health" { render_ok() } else { @@ -68,7 +75,7 @@ pub fn health(req: Request) -> RoutingResult { } /// Produce prometheus style metrics -pub fn metrics(req: Request) -> RoutingResult { +pub async fn metrics(req: Request) -> RoutingResult { if req.is_get() && req.path().segment().starts_with("metrics") { let server = req.read(); @@ -181,7 +188,7 @@ pub fn metrics(req: Request) -> RoutingResult { } /// Return various stats as json -pub fn stats(req: Request) -> RoutingResult { +pub async fn stats(req: Request) -> RoutingResult { if !req.is_get() { Err(req) } else if req.path().full() == "/stats/info" { @@ -196,7 +203,7 @@ pub fn stats(req: Request) -> RoutingResult { } /// Maps the API methods -pub fn api(mut req: Request) -> RoutingResult { +pub async fn api(req: Request) -> RoutingResult { if !req.path().full().starts_with("/api/v1") { Err(req) // Not for us } else { @@ -211,13 +218,13 @@ pub fn api(mut req: Request) -> RoutingResult { match path.next() { Some("authorized") => api_authorized(req), - Some("publishers") => api_publishers(req, &mut path), + Some("publishers") => api_publishers(req, &mut path).await, _ => render_unknown_method(), } } } -fn api_authorized(mut req: Request) -> RoutingResult { +fn api_authorized(req: Request) -> RoutingResult { if req.is_get() { render_ok() } else { @@ -225,7 +232,7 @@ fn api_authorized(mut req: Request) -> RoutingResult { } } -fn api_publishers(mut req: Request, path: &mut RequestPath) -> RoutingResult { +async fn api_publishers(req: Request, path: &mut RequestPath) -> RoutingResult { if req.is_get() { if let Some(publisher_str) = path.next() { let publisher = match PublisherHandle::from_str(publisher_str) { @@ -244,7 +251,10 @@ fn api_publishers(mut req: Request, path: &mut RequestPath) -> RoutingResult { list_pbl(req) } } else if req.is_post() { - unimplemented!("Get post body") + match path.next() { + None => add_pbl(req).await, + _ => render_unknown_method(), + } } else { render_unknown_method() } @@ -274,8 +284,13 @@ pub fn list_pbl(req: Request) -> RoutingResult { } /// Adds a publisher -pub fn add_pbl(req: Request, pbl: rfc8183::PublisherRequest) -> RoutingResult { - render_json_res(req.write().add_publisher(pbl)) +async fn add_pbl(req: Request) -> RoutingResult { + let server = req.state().clone(); + render_json_res_res( + req.json() + .await + .map(|pbl| server.write().add_publisher(pbl)), + ) } /// Removes a publisher. Should be idempotent! If if did not exist then diff --git a/src/daemon/http/mod.rs b/src/daemon/http/mod.rs index 42320062..dc22c9f9 100644 --- a/src/daemon/http/mod.rs +++ b/src/daemon/http/mod.rs @@ -1,8 +1,12 @@ -use std::io; use std::sync::{RwLockReadGuard, RwLockWriteGuard}; +use std::{fmt, io}; +use serde::de::DeserializeOwned; use serde::Serialize; +use bytes::{Buf, BufMut, Bytes}; + +use hyper::body::HttpBody; use hyper::http::uri::PathAndQuery; use hyper::{Body, Method, StatusCode}; @@ -162,13 +166,13 @@ impl HttpResponse { //------------ Request ------------------------------------------------------- pub struct Request { - request: hyper::Request, + request: hyper::Request, path: RequestPath, state: State, } impl Request { - pub fn new(request: hyper::Request, state: State) -> Self { + pub fn new(request: hyper::Request, state: State) -> Self { let path = RequestPath::from_request(&request); Request { request, @@ -183,7 +187,7 @@ impl Request { } /// Get the application State - fn state(&self) -> &State { + pub fn state(&self) -> &State { &self.state } @@ -212,6 +216,79 @@ impl Request { self.request.method() == Method::POST } + /// Get a json object from a post body + pub async fn json(mut self) -> Result { + let limit = self.read().limit_api(); + let body = self.request.into_body(); + + let bytes = Self::to_bytes_limited(body, limit) + .await + .map_err(|_| Error::custom("Error reading body"))?; + serde_json::from_slice(&bytes).map_err(Error::JsonError) + } + + /// See hyper::body::to_bytes + /// + /// Here we want to limit the bytes consumed to a maximum. So, the + /// code below is adapted from the method in the hyper crate. + async fn to_bytes_limited(body: T, limit: usize) -> Result + where + T: HttpBody, + { + futures_util::pin_mut!(body); + + let mut size_processed = 0; + + fn assert_body_size(size: usize, limit: usize) -> Result<(), io::Error> { + if size > limit { + Err(io::Error::new( + io::ErrorKind::Other, + "Post exceeds max length", + )) + } else { + Ok(()) + } + } + + // If there's only 1 chunk, we can just return Buf::to_bytes() + let mut first = if let Some(buf) = body.data().await { + let buf = buf.map_err(|_| RequestError::Hyper)?; + let size = buf.bytes().len(); + size_processed += size; + assert_body_size(size_processed, limit)?; + buf + } else { + return Ok(Bytes::new()); + }; + + let second = if let Some(buf) = body.data().await { + let buf = buf.map_err(|_| RequestError::Hyper)?; + let size = buf.bytes().len(); + size_processed += size; + assert_body_size(size_processed, limit)?; + buf + } else { + return Ok(first.to_bytes()); + }; + + // With more than 1 buf, we gotta flatten into a Vec first. + let cap = first.remaining() + second.remaining() + body.size_hint().lower() as usize; + let mut vec = Vec::with_capacity(cap); + vec.put(first); + vec.put(second); + + while let Some(buf) = body.data().await { + let buf = buf.map_err(|_| RequestError::Hyper)?; + let size = buf.bytes().len(); + size_processed += size; + assert_body_size(size_processed, limit)?; + + vec.put(buf); + } + + Ok(vec.into()) + } + /// Checks whether the Bearer token is set to what we expect pub fn is_authorized(&self) -> bool { if let Some(header) = self.request.headers().get("Authorization") { @@ -231,6 +308,17 @@ impl Request { } } +pub enum RequestError { + Hyper, + Io(io::Error), +} + +impl From for RequestError { + fn from(e: io::Error) -> Self { + RequestError::Io(e) + } +} + //------------ RequestPath --------------------------------------------------- #[derive(Clone)] diff --git a/src/daemon/http/server.rs b/src/daemon/http/server.rs index 122df576..0548fe8b 100644 --- a/src/daemon/http/server.rs +++ b/src/daemon/http/server.rs @@ -44,7 +44,7 @@ pub async fn start(config: Config) -> Result<(), Error> { async move { Ok::<_, Infallible>(service_fn(move |req: hyper::Request| { let mut state = state.clone(); - async move { map_requests(req, state) } + map_requests(req, state) })) } }); @@ -74,7 +74,7 @@ pub async fn start(config: Config) -> Result<(), Error> { Ok(()) } -fn map_requests( +async fn map_requests( req: hyper::Request, mut state: State, ) -> Result, Error> { @@ -87,7 +87,8 @@ fn map_requests( .or_else(stats) .or_else(api) .or_else(render_not_found) - .map_err(|_| Error::custom("should have received not found response"))? + .map_err(|_| Error::custom("should have received not found response")) + .await? .res() // // let post_limit_api = config.post_limit_api; diff --git a/src/daemon/krillserver.rs b/src/daemon/krillserver.rs index f14d8966..7d12f470 100644 --- a/src/daemon/krillserver.rs +++ b/src/daemon/krillserver.rs @@ -34,17 +34,7 @@ use crate::publish::CaPublisher; //------------ KrillServer --------------------------------------------------- /// This is the master krill server that is doing all the orchestration -/// for all the components, like: -/// * Admin tasks: -/// * Verify (admin) API access -/// * Manage known publishers -/// * CMS proxy: -/// * Decodes and validates CMS sent by known publishers using CMS -/// * Encodes and signs CMS responses for remote publishers using CMS -/// * Repository: -/// * Process publish / list requests by known publishers -/// * Updates the repository on disk -/// * Updates the RRDP files +/// for all the components. pub struct KrillServer { // The base URI for this service service_uri: uri::Https, @@ -67,6 +57,35 @@ pub struct KrillServer { // Time this server was started started: Time, + + // Global size constraints on things which can be posted + post_limits: PostLimits, +} + +pub struct PostLimits { + api: usize, + rfc6492: usize, + rfc8181: usize, +} + +impl PostLimits { + fn new(api: usize, rfc6492: usize, rfc8181: usize) -> Self { + PostLimits { + api, + rfc8181, + rfc6492, + } + } + + pub fn api(&self) -> usize { + self.api + } + pub fn rfc6492(&self) -> usize { + self.rfc6492 + } + pub fn rfc8181(&self) -> usize { + self.rfc8181 + } } /// # Set up and initialisation @@ -152,6 +171,12 @@ impl KrillServer { ca_refresh_rate, ); + let post_limits = PostLimits::new( + config.post_limit_api, + config.post_limit_rfc6492, + config.post_limit_rfc8181, + ); + Ok(KrillServer { service_uri, work_dir: work_dir.clone(), @@ -160,6 +185,7 @@ impl KrillServer { caserver, scheduler, started: Time::now(), + post_limits, }) } @@ -172,11 +198,15 @@ impl KrillServer { } } -/// # Authentication +/// # Authentication and Access impl KrillServer { pub fn is_api_allowed(&self, auth: &Auth) -> bool { self.authorizer.is_api_allowed(auth) } + + pub fn limit_api(&self) -> usize { + self.post_limits.api() + } } /// # Configure publishers diff --git a/src/lib.rs b/src/lib.rs index 53ce6d29..7481a616 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -8,6 +8,7 @@ extern crate clokwerk; #[macro_use] extern crate derive_more; extern crate futures; +extern crate futures_util; extern crate hex; extern crate hyper; #[macro_use]