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
23 changes: 22 additions & 1 deletion bd-grpc-codec/src/code.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
}
}
Expand All @@ -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,
}
Expand Down
29 changes: 29 additions & 0 deletions bd-grpc-codec/src/coding_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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,
Expand Down
11 changes: 9 additions & 2 deletions bd-grpc/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -235,20 +235,27 @@ impl<C: Connect + Clone + Send + Sync + 'static> Client<C> {
},
};
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(),
);
}

// 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)
Expand Down
7 changes: 5 additions & 2 deletions bd-grpc/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -531,11 +531,14 @@ impl<IncomingType: DecodingResult> StreamingApiReceiver<IncomingType> {
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));
}

Expand Down
46 changes: 43 additions & 3 deletions bd-grpc/src/status.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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",
}
Expand Down Expand Up @@ -125,6 +137,8 @@ pub struct Status {
code: Code,
message: Option<String>,
original_error: Option<String>,
response_status: Option<StatusCode>,
response_headers: Option<Arc<HeaderMap>>,
}

impl PartialEq for Status {
Expand All @@ -150,6 +164,8 @@ impl Status {
code,
message: Some(message.into()),
original_error,
response_status: None,
response_headers: None,
}
}

Expand All @@ -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<StatusCode> {
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<T>(response: &http::Response<T>) -> Option<&str> {
response
Expand All @@ -196,6 +234,8 @@ impl Status {
code,
message,
original_error: None,
response_status: None,
response_headers: None,
}
}

Expand Down
54 changes: 53 additions & 1 deletion bd-grpc/src/status_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand All @@ -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
);
}
}
Loading