use axum::extract::rejection::JsonRejection;
use toolkit_canonical_errors::{CanonicalError, Http, resource_error};
#[resource_error(gts_id!("cf.core.http.request.v1~"))]
pub struct GenericResourceError;
pub fn json_rejection_code(rejection: &JsonRejection) -> &'static str {
match rejection {
JsonRejection::JsonSyntaxError(_) => "json_syntax_error",
JsonRejection::JsonDataError(_) => "invalid_json_body",
JsonRejection::MissingJsonContentType(_) => "missing_json_content_type",
JsonRejection::BytesRejection(_) => "json_body_read_error",
_ => {
debug_assert!(
false,
"unhandled JsonRejection variant, update json_rejection_code: {rejection}"
);
tracing::error!(
rejection = %rejection,
"extract::Json: unhandled JsonRejection variant, update json_rejection_code"
);
"unclassified_json_rejection"
}
}
}
pub fn json_rejection_to_canonical(rejection: &JsonRejection) -> CanonicalError {
let code = json_rejection_code(rejection);
rejection_to_canonical(
"body",
code,
rejection.status().as_u16(),
rejection.body_text(),
)
}
const MAX_FIELD_VIOLATION_DESCRIPTION_CHARS: usize = 500;
pub fn rejection_to_canonical(
field: &str,
code: &str,
status: u16,
message: String,
) -> CanonicalError {
if status >= 500 {
return if status == 500 {
CanonicalError::internal(message).create()
} else {
CanonicalError::internal(message)
.with_override(Http::status_code(status))
.create()
};
}
GenericResourceError::invalid_argument()
.with_field_violation(field, truncate_description(message), code)
.with_override(Http::status_code(status))
.create()
}
const TRUNCATION_SUFFIX: &str = "... (truncated)";
fn truncate_description(message: String) -> String {
if message.chars().count() <= MAX_FIELD_VIOLATION_DESCRIPTION_CHARS {
return message;
}
tracing::debug!(
original_len = message.len(),
"extractor rejection message exceeded MAX_FIELD_VIOLATION_DESCRIPTION_CHARS, truncating before sending to client"
);
let mut truncated: String = message
.chars()
.take(MAX_FIELD_VIOLATION_DESCRIPTION_CHARS - TRUNCATION_SUFFIX.chars().count())
.collect();
truncated.push_str(TRUNCATION_SUFFIX);
truncated
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests {
use super::*;
#[test]
fn short_message_passes_through_unchanged() {
assert_eq!(
truncate_description("invalid digit found in string".to_owned()),
"invalid digit found in string"
);
}
#[test]
fn rejection_to_canonical_maps_any_5xx_status_to_internal_not_invalid_argument() {
let err = rejection_to_canonical("body", "some_4xx_only_code", 503, "boom".to_owned());
let problem: toolkit_canonical_errors::Problem = err.into();
let json = serde_json::to_value(&problem).unwrap();
assert_eq!(
json,
serde_json::json!({
"type": "gts://gts.cf.core.errors.err.v1~cf.core.err.internal.v1~",
"title": "Internal",
"status": 503,
"detail": "An internal error occurred. Please retry later.",
"context": {},
})
);
}
#[test]
fn oversized_message_is_truncated_with_a_marker() {
let huge = "x".repeat(MAX_FIELD_VIOLATION_DESCRIPTION_CHARS + 1000);
let result = truncate_description(huge);
assert_eq!(
result.chars().count(),
MAX_FIELD_VIOLATION_DESCRIPTION_CHARS
);
assert_eq!(
result,
format!(
"{}... (truncated)",
"x".repeat(
MAX_FIELD_VIOLATION_DESCRIPTION_CHARS - TRUNCATION_SUFFIX.chars().count()
)
)
);
}
}