Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions aw-server/src/endpoints/bucket.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ use rocket::http::Status;
use rocket::State;

use crate::endpoints::query_cache::event_range;
use crate::endpoints::util::{BucketEventsCsvRocket, BucketsExportRocket};
use crate::endpoints::util::{ApiJson, BucketEventsCsvRocket, BucketsExportRocket};
use crate::endpoints::{HttpErrorJson, ServerState};

#[get("/")]
Expand Down Expand Up @@ -48,7 +48,7 @@ pub fn bucket_get(
#[post("/<bucket_id>", data = "<message>", format = "application/json")]
pub fn bucket_new(
bucket_id: &str,
message: Json<Bucket>,
message: ApiJson<Bucket>,
state: &State<ServerState>,
) -> Result<(), HttpErrorJson> {
let mut bucket = message.into_inner();
Expand Down Expand Up @@ -170,7 +170,7 @@ pub fn bucket_events_get_single(
#[post("/<bucket_id>/events", data = "<events>", format = "application/json")]
pub fn bucket_events_create(
bucket_id: &str,
events: Json<Vec<Event>>,
events: ApiJson<Vec<Event>>,
state: &State<ServerState>,
) -> Result<Json<Vec<Event>>, HttpErrorJson> {
// Hold the write lock across (read old ranges + write + invalidate) so a
Expand Down Expand Up @@ -216,7 +216,7 @@ pub fn bucket_events_create(
)]
pub fn bucket_events_heartbeat(
bucket_id: &str,
heartbeat_json: Json<Event>,
heartbeat_json: ApiJson<Event>,
pulsetime: f64,
state: &State<ServerState>,
) -> Result<Json<Event>, HttpErrorJson> {
Expand Down
3 changes: 2 additions & 1 deletion aw-server/src/endpoints/import.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ use aw_models::{BucketsExport, Event, TryVec};

use aw_datastore::{Datastore, DatastoreError};

use crate::endpoints::util::ApiJson;
use crate::endpoints::{HttpErrorJson, ServerState};

/// Computes a dedup identity tuple for an event.
Expand Down Expand Up @@ -132,7 +133,7 @@ fn import(datastore: &Datastore, import: BucketsExport) -> Result<(), HttpErrorJ
#[post("/", data = "<json_data>", format = "application/json")]
pub fn bucket_import_json(
state: &State<ServerState>,
json_data: Json<BucketsExport>,
json_data: ApiJson<BucketsExport>,
) -> Result<(), HttpErrorJson> {
let result = import(&state.datastore, json_data.into_inner());
// Clear even on failure: a multi-bucket import can write earlier buckets
Expand Down
22 changes: 21 additions & 1 deletion aw-server/src/endpoints/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ use std::path::{Path, PathBuf};
use gethostname::gethostname;
use rocket::fs::FileServer;
use rocket::http::{ContentType, Status};
use rocket::request::Request;
use rocket::serde::json::Json;
use rocket::State;

Expand Down Expand Up @@ -271,7 +272,11 @@ pub fn build_rocket(server_state: ServerState, config: AWConfig) -> rocket::Rock
settings::settings_get,
],
)
.mount("/", rocket_cors::catch_all_options_routes());
.mount("/", rocket_cors::catch_all_options_routes())
.register(
"/api",
catchers![api_bad_request, api_not_found, api_unprocessable],
);

// for each custom static directory, mount it at the given name
for (name, dir) in custom_static {
Expand All @@ -284,6 +289,21 @@ pub fn build_rocket(server_state: ServerState, config: AWConfig) -> rocket::Rock
rocket
}

#[catch(400)]
fn api_bad_request(req: &Request) -> util::HttpErrorJson {
util::api_error_json(Status::BadRequest, req)
}

#[catch(404)]
fn api_not_found(req: &Request) -> util::HttpErrorJson {
util::api_error_json(Status::NotFound, req)
}

#[catch(422)]
fn api_unprocessable(req: &Request) -> util::HttpErrorJson {
util::api_error_json(Status::UnprocessableEntity, req)
}

mod tests {
#[test]
fn test_filesystem_resolver() {
Expand Down
4 changes: 2 additions & 2 deletions aw-server/src/endpoints/query.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,13 @@ use std::sync::Arc;

use rocket::http::Status;
use rocket::response::content::RawJson;
use rocket::serde::json::Json;
use rocket::State;

use aw_models::Query;
use aw_query::QueryError;

use crate::endpoints::query_cache::CacheKey;
use crate::endpoints::util::ApiJson;
use crate::endpoints::{HttpErrorJson, ServerState};

fn query_error_status(e: &QueryError) -> Status {
Expand Down Expand Up @@ -59,7 +59,7 @@ mod tests {
/// cache miss does not pay a second full serialization for the response.
#[post("/?<cache>", data = "<query_req>", format = "application/json")]
pub fn query(
query_req: Json<Query>,
query_req: ApiJson<Query>,
cache: Option<bool>,
state: &State<ServerState>,
) -> Result<RawJson<String>, HttpErrorJson> {
Expand Down
3 changes: 2 additions & 1 deletion aw-server/src/endpoints/settings.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ use std::collections::HashMap;

use aw_datastore::DatastoreError;

use crate::endpoints::util::ApiJson;
use crate::endpoints::HttpErrorJson;

/// Map a settings API key to the datastore key (`settings.<key>`).
Expand Down Expand Up @@ -107,7 +108,7 @@ pub fn setting_get(
pub fn setting_set(
state: &State<ServerState>,
key: String,
value: Json<serde_json::Value>,
value: ApiJson<serde_json::Value>,
) -> Result<Status, HttpErrorJson> {
let setting_key = parse_key(key)?;
let value_str = match serde_json::to_string(&value.0) {
Expand Down
52 changes: 52 additions & 0 deletions aw-server/src/endpoints/util.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,14 @@ use std::io::{copy, pipe, Cursor, PipeReader, PipeWriter, Seek, SeekFrom};
use std::thread;

use chrono::{DateTime, Utc};
use rocket::data::{self, Data, FromData};
use rocket::http::ContentType;
use rocket::http::Header;
use rocket::http::Status;
use rocket::request::Request;
use rocket::response::{self, Responder, Response};
use rocket::serde::json::Json;
use serde::Deserialize;
use serde::Serialize;

#[derive(Serialize, Debug)]
Expand Down Expand Up @@ -40,6 +43,55 @@ impl<'r> Responder<'r, 'static> for HttpErrorJson {
}
}

/// Reason a JSON request body was rejected, stashed in the request-local cache
/// so the `/api` error catchers can report it (Rocket only logs it).
#[derive(Default)]
struct BodyError(std::sync::Mutex<Option<String>>);

/// Drop-in replacement for `Json<T>` as a data guard: identical parsing, but the
/// serde error is kept so [`api_error_json`] can return it to the client.
pub struct ApiJson<T>(pub T);

impl<T> ApiJson<T> {
pub fn into_inner(self) -> T {
self.0
}
}

impl<T> std::ops::Deref for ApiJson<T> {
type Target = T;
fn deref(&self) -> &T {
&self.0
}
}

#[rocket::async_trait]
impl<'r, T: Deserialize<'r>> FromData<'r> for ApiJson<T> {
type Error = String;

async fn from_data(req: &'r Request<'_>, data: Data<'r>) -> data::Outcome<'r, Self> {
match <Json<T> as FromData>::from_data(req, data).await {
data::Outcome::Success(json) => data::Outcome::Success(ApiJson(json.into_inner())),
data::Outcome::Error((status, err)) => {
let msg = err.to_string();
*req.local_cache(BodyError::default).0.lock().unwrap() = Some(msg.clone());
data::Outcome::Error((status, msg))
}
data::Outcome::Forward(data) => data::Outcome::Forward(data),
}
}
}

/// Build the JSON error body shared by the `/api` catchers.
pub fn api_error_json(status: Status, req: &Request) -> HttpErrorJson {
let reason = req.local_cache(BodyError::default).0.lock().unwrap().take();
let message = match reason {
Some(detail) => format!("{}: {}", status.reason_lossy(), detail),
None => status.reason_lossy().to_string(),
};
HttpErrorJson::new(status, message)
}

pub struct BucketsExportRocket {
datastore: aw_datastore::Datastore,
bucket_id: Option<String>,
Expand Down
50 changes: 50 additions & 0 deletions aw-server/tests/api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1003,6 +1003,56 @@ mod api_tests {
res.status()
}

#[test]
Comment thread
TimeToBuildBob marked this conversation as resolved.
fn test_api_errors_are_json_with_reason() {
let server = setup_testserver();
let client = Client::untracked(server).unwrap();

// Malformed event body -> 422 JSON carrying the serde error
let res = client
.post("/api/0/buckets/b1/events")
.header(ContentType::JSON)
.header(Header::new("Host", "127.0.0.1:5600"))
.body("{\"duration\": 0}")
.dispatch();
assert_eq!(res.status(), Status::UnprocessableEntity);
assert_eq!(res.content_type(), Some(ContentType::JSON));
let body: Value = serde_json::from_str(&res.into_string().unwrap()).unwrap();
let msg = body["message"].as_str().unwrap();
assert!(msg.starts_with("Unprocessable Entity: "), "{msg}");
assert!(msg.contains("line 1"), "{msg}");

// Syntactically invalid JSON -> JSON too
let res = client
.post("/api/0/buckets/b1/events")
.header(ContentType::JSON)
.header(Header::new("Host", "127.0.0.1:5600"))
.body("{not json")
.dispatch();
assert_eq!(res.status(), Status::UnprocessableEntity);
assert_eq!(res.content_type(), Some(ContentType::JSON));

// Unknown API route -> 404 JSON
let res = client
.get("/api/0/nope")
.header(Header::new("Host", "127.0.0.1:5600"))
.dispatch();
assert_eq!(res.status(), Status::NotFound);
assert_eq!(res.content_type(), Some(ContentType::JSON));
let body: Value = serde_json::from_str(&res.into_string().unwrap()).unwrap();
assert_eq!(body["message"], "Not Found");

// Rocket-generated 400 (invalid request URI) -> JSON from the 400 catcher
let res = client
.get("/api/0/bad path")
.header(Header::new("Host", "127.0.0.1:5600"))
.dispatch();
assert_eq!(res.status(), Status::BadRequest);
assert_eq!(res.content_type(), Some(ContentType::JSON));
let body: Value = serde_json::from_str(&res.into_string().unwrap()).unwrap();
assert_eq!(body["message"], "Bad Request");
}

#[test]
fn test_illegally_long_key() {
let server = setup_testserver();
Expand Down
Loading