diff --git a/bd-grpc-codec/src/code.rs b/bd-grpc-codec/src/code.rs index 4a93b8a50..155058955 100644 --- a/bd-grpc-codec/src/code.rs +++ b/bd-grpc-codec/src/code.rs @@ -8,15 +8,22 @@ // Code // -// Wrapper for supported gRPC status codes. Unknown is a synthetic code if mapping is not possible. +// Wrapper for gRPC status codes. Unknown is the fallback for invalid wire values. #[derive(PartialEq, Eq, Debug, Clone, Copy)] pub enum Code { Ok, + Cancelled, Unknown, InvalidArgument, + DeadlineExceeded, + AlreadyExists, FailedPrecondition, + Aborted, + OutOfRange, + Unimplemented, Internal, Unavailable, + DataLoss, Unauthenticated, NotFound, PermissionDenied, @@ -29,14 +36,21 @@ impl Code { pub const fn to_int(&self) -> i32 { match self { Self::Ok => 0, + Self::Cancelled => 1, Self::Unknown => 2, Self::InvalidArgument => 3, + Self::DeadlineExceeded => 4, Self::NotFound => 5, + Self::AlreadyExists => 6, Self::PermissionDenied => 7, Self::ResourceExhausted => 8, Self::FailedPrecondition => 9, + Self::Aborted => 10, + Self::OutOfRange => 11, + Self::Unimplemented => 12, Self::Internal => 13, Self::Unavailable => 14, + Self::DataLoss => 15, Self::Unauthenticated => 16, } } @@ -47,13 +61,20 @@ impl Code { pub fn from_str(status: &str) -> Self { match status { "0" => Self::Ok, + "1" => Self::Cancelled, "3" => Self::InvalidArgument, + "4" => Self::DeadlineExceeded, "5" => Self::NotFound, + "6" => Self::AlreadyExists, "7" => Self::PermissionDenied, "8" => Self::ResourceExhausted, "9" => Self::FailedPrecondition, + "10" => Self::Aborted, + "11" => Self::OutOfRange, + "12" => Self::Unimplemented, "13" => Self::Internal, "14" => Self::Unavailable, + "15" => Self::DataLoss, "16" => Self::Unauthenticated, _ => Self::Unknown, } diff --git a/bd-grpc-codec/src/coding_test.rs b/bd-grpc-codec/src/coding_test.rs index c085e6bfb..ab6b213e0 100644 --- a/bd-grpc-codec/src/coding_test.rs +++ b/bd-grpc-codec/src/coding_test.rs @@ -4,6 +4,7 @@ // Licensed under the Apache License, Version 2.0 (the "License"); // you may not use this file except in compliance with the License. +use crate::code::Code; use crate::{Compression, DEFAULT_MAX_MESSAGE_BYTES, Decoder, Decompression, Encoder, OptimizeFor}; use protobuf::Message; use protobuf::well_known_types::any::Any; @@ -15,6 +16,34 @@ fn test_global_init() { bd_test_helpers_core::test_global_init(); } +#[rstest] +#[case("0", Code::Ok)] +#[case("1", Code::Cancelled)] +#[case("2", Code::Unknown)] +#[case("3", Code::InvalidArgument)] +#[case("4", Code::DeadlineExceeded)] +#[case("5", Code::NotFound)] +#[case("6", Code::AlreadyExists)] +#[case("7", Code::PermissionDenied)] +#[case("8", Code::ResourceExhausted)] +#[case("9", Code::FailedPrecondition)] +#[case("10", Code::Aborted)] +#[case("11", Code::OutOfRange)] +#[case("12", Code::Unimplemented)] +#[case("13", Code::Internal)] +#[case("14", Code::Unavailable)] +#[case("15", Code::DataLoss)] +#[case("16", Code::Unauthenticated)] +fn grpc_status_codes_round_trip(#[case] wire_code: &str, #[case] expected: Code) { + assert_eq!(Code::from_str(wire_code), expected); + assert_eq!(expected.to_int().to_string(), wire_code); +} + +#[test] +fn invalid_grpc_status_code_maps_to_unknown() { + assert_eq!(Code::from_str("17"), Code::Unknown); +} + #[rstest] #[case((Compression::StatefulZlib { level: 3, diff --git a/bd-grpc/src/client.rs b/bd-grpc/src/client.rs index 16a8942a8..d17260fbf 100644 --- a/bd-grpc/src/client.rs +++ b/bd-grpc/src/client.rs @@ -235,12 +235,15 @@ impl Client { }, }; if !response.status().is_success() { + let response_status = response.status(); + let response_headers = response.headers().clone(); return Err( Status::new( Code::Internal, - format!("Non-200 response code: {}", response.status()), + format!("Non-200 response code: {response_status}"), None, ) + .with_response_context(response_status, response_headers) .into(), ); } @@ -248,7 +251,11 @@ impl Client { // We treat any trailer only response as an error, even with the response status is OK. This // seems fine for now. if response.headers().contains_key(GRPC_STATUS) { - return Err(Status::from_headers(response.headers()).into()); + return Err( + Status::from_headers(response.headers()) + .with_response_context(response.status(), response.headers().clone()) + .into(), + ); } Ok(response) diff --git a/bd-grpc/src/lib.rs b/bd-grpc/src/lib.rs index 10da27f64..a158a3bd8 100644 --- a/bd-grpc/src/lib.rs +++ b/bd-grpc/src/lib.rs @@ -47,7 +47,7 @@ use bd_stats_common::DynCounter; use bytes::{BufMut, Bytes, BytesMut}; use connect_protocol::{ConnectProtocolType, EndOfStreamResponse, ErrorResponse, ToContentType}; use http::header::{CONTENT_ENCODING, CONTENT_TYPE}; -use http::{Extensions, HeaderMap}; +use http::{Extensions, HeaderMap, StatusCode}; use http_body::Frame; use http_body_util::{BodyExt, LengthLimitError, StreamBody}; use protobuf::{Message, MessageFull}; @@ -531,11 +531,14 @@ impl StreamingApiReceiver { return Ok(None); } + let mut response_headers = self.headers.clone(); + response_headers.extend(trailers.clone()); let status = Status::from_wire( code, grpc_message .map(|value| Status::decode_grpc_message(value.to_str().unwrap_or_default())), - ); + ) + .with_response_context(StatusCode::OK, response_headers); return Err(Error::Grpc(status)); } diff --git a/bd-grpc/src/status.rs b/bd-grpc/src/status.rs index 03db36f23..46a7a675c 100644 --- a/bd-grpc/src/status.rs +++ b/bd-grpc/src/status.rs @@ -15,14 +15,19 @@ use axum::response::Response; use bd_grpc_codec::code::Code; use http::header::CONTENT_TYPE; use http::{Extensions, HeaderMap, HeaderValue, StatusCode}; +use std::sync::Arc; // https://connectrpc.com/docs/protocol#error-codes #[must_use] -pub const fn code_to_connect_http_status(code: Code) -> StatusCode { +pub fn code_to_connect_http_status(code: Code) -> StatusCode { match code { Code::Ok => StatusCode::OK, - Code::Unknown | Code::Internal => StatusCode::INTERNAL_SERVER_ERROR, - Code::InvalidArgument | Code::FailedPrecondition => StatusCode::BAD_REQUEST, + Code::Cancelled => StatusCode::from_u16(499).unwrap_or(StatusCode::BAD_REQUEST), + Code::Unknown | Code::Internal | Code::DataLoss => StatusCode::INTERNAL_SERVER_ERROR, + Code::InvalidArgument | Code::FailedPrecondition | Code::OutOfRange => StatusCode::BAD_REQUEST, + Code::DeadlineExceeded => StatusCode::GATEWAY_TIMEOUT, + Code::AlreadyExists | Code::Aborted => StatusCode::CONFLICT, + Code::Unimplemented => StatusCode::NOT_IMPLEMENTED, Code::Unavailable => StatusCode::SERVICE_UNAVAILABLE, Code::Unauthenticated => StatusCode::UNAUTHORIZED, Code::NotFound => StatusCode::NOT_FOUND, @@ -36,13 +41,20 @@ pub const fn code_to_connect_http_status(code: Code) -> StatusCode { pub const fn code_to_connect_code_string(code: Code) -> &'static str { match code { Code::Ok => "ok", + Code::Cancelled => "canceled", Code::Unknown => "unknown", Code::InvalidArgument => "invalid_argument", + Code::DeadlineExceeded => "deadline_exceeded", Code::FailedPrecondition => "failed_precondition", + Code::Aborted => "aborted", + Code::OutOfRange => "out_of_range", + Code::Unimplemented => "unimplemented", Code::Internal => "internal", Code::Unavailable => "unavailable", + Code::DataLoss => "data_loss", Code::Unauthenticated => "unauthenticated", Code::NotFound => "not_found", + Code::AlreadyExists => "already_exists", Code::PermissionDenied => "permission_denied", Code::ResourceExhausted => "resource_exhausted", } @@ -125,6 +137,8 @@ pub struct Status { code: Code, message: Option, original_error: Option, + response_status: Option, + response_headers: Option>, } impl PartialEq for Status { @@ -150,6 +164,8 @@ impl Status { code, message: Some(message.into()), original_error, + response_status: None, + response_headers: None, } } @@ -173,6 +189,28 @@ impl Status { self.original_error_message().or_else(|| self.message()) } + /// Attaches the HTTP response metadata that accompanied this status. + #[must_use] + pub fn with_response_context( + mut self, + response_status: StatusCode, + response_headers: HeaderMap, + ) -> Self { + self.response_status = Some(response_status); + self.response_headers = Some(Arc::new(response_headers)); + self + } + + #[must_use] + pub const fn response_status(&self) -> Option { + self.response_status + } + + #[must_use] + pub fn response_headers(&self) -> Option<&HeaderMap> { + self.response_headers.as_deref() + } + #[must_use] pub fn trace_error_message_from_response(response: &http::Response) -> Option<&str> { response @@ -196,6 +234,8 @@ impl Status { code, message, original_error: None, + response_status: None, + response_headers: None, } } diff --git a/bd-grpc/src/status_test.rs b/bd-grpc/src/status_test.rs index 460589938..d3e4f2991 100644 --- a/bd-grpc/src/status_test.rs +++ b/bd-grpc/src/status_test.rs @@ -4,9 +4,12 @@ // Licensed under the Apache License, Version 2.0 (the "License"); // you may not use this file except in compliance with the License. -use super::Status; +use super::{Status, code_to_connect_code_string, code_to_connect_http_status}; use axum::body::Body; use axum::response::Response; +use bd_grpc_codec::code::Code; +use http::header::HeaderName; +use http::{HeaderMap, HeaderValue, StatusCode}; #[test] fn set_trace_error_message_attaches_message_to_response() { @@ -22,3 +25,52 @@ fn set_trace_error_message_attaches_message_to_response() { Some("original error") ); } + +#[test] +fn response_context_preserves_status_and_headers() { + let response_headers = HeaderMap::from_iter([( + HeaderName::from_static("x-request-id"), + HeaderValue::from_static("request-123"), + )]); + let status = Status::new(Code::Internal, "upstream error", None) + .with_response_context(StatusCode::BAD_GATEWAY, response_headers); + + assert_eq!(status.response_status(), Some(StatusCode::BAD_GATEWAY)); + assert_eq!( + status + .response_headers() + .and_then(|headers| headers.get("x-request-id")), + Some(&HeaderValue::from_static("request-123")) + ); +} + +#[test] +fn all_grpc_codes_have_connect_mappings() { + let mappings = [ + (Code::Ok, "ok", 200), + (Code::Cancelled, "canceled", 499), + (Code::Unknown, "unknown", 500), + (Code::InvalidArgument, "invalid_argument", 400), + (Code::DeadlineExceeded, "deadline_exceeded", 504), + (Code::NotFound, "not_found", 404), + (Code::AlreadyExists, "already_exists", 409), + (Code::PermissionDenied, "permission_denied", 403), + (Code::ResourceExhausted, "resource_exhausted", 429), + (Code::FailedPrecondition, "failed_precondition", 400), + (Code::Aborted, "aborted", 409), + (Code::OutOfRange, "out_of_range", 400), + (Code::Unimplemented, "unimplemented", 501), + (Code::Internal, "internal", 500), + (Code::Unavailable, "unavailable", 503), + (Code::DataLoss, "data_loss", 500), + (Code::Unauthenticated, "unauthenticated", 401), + ]; + + for (code, expected_connect_code, expected_http_status) in mappings { + assert_eq!(code_to_connect_code_string(code), expected_connect_code); + assert_eq!( + code_to_connect_http_status(code).as_u16(), + expected_http_status + ); + } +}