feat(tonic): add h2::Error as a source for Status (#612)

## Motivation

A gRPC server may send a HTTP/2 GOAWAY frame with NO_ERROR status to gracefully shutdown a connection. This appears to Tonic users as a `tonic::Status` with `Code::Internal` and the message set to `h2 protocol error: protocol error: not a result of an error`.

The only way to currently detect this case and differentiate it from other internal errors (e.g., an application-level internal error) is to match on the message. A client may want to differentiate these cases because it may only want to alert on the application-level internal error and not on the transient transport-level issue. (Indeed, this is the use case for which I'm envisioning using this change.)

Matching on a message is not as robust, however, as matching on an `h2::Error` and its reason code. (The message could change for example if a future version of Tonic decided to vary the message. This would break any users that matched on the previous version of the message.)

## Solution

Store the `h2::Error` used when creating a `tonic::Status` from a `h2::Error` and provide it as the `source` for purposes of `std::error::Error`. This will allow users to downcast it and match on the original `h2::Reason`.
This commit is contained in:
Tom Dyas
2021-06-23 08:34:07 +02:00
committed by GitHub
parent 12815d0a1d
commit b90bb7bbc0
5 changed files with 93 additions and 46 deletions
+1 -1
View File
@@ -164,7 +164,7 @@ impl<T> Grpc<T> {
.inner .inner
.call(request) .call(request)
.await .await
.map_err(|err| Status::from_error(&*(err.into())))?; .map_err(|err| Status::from_error(err.into()))?;
let status_code = response.status(); let status_code = response.status();
let trailers_only_status = Status::from_header_map(response.headers()); let trailers_only_status = Status::from_header_map(response.headers());
+4 -4
View File
@@ -149,9 +149,9 @@ impl<T> Streaming<T> {
// them manually. // them manually.
let map = future::poll_fn(|cx| Pin::new(&mut self.body).poll_trailers(cx)) let map = future::poll_fn(|cx| Pin::new(&mut self.body).poll_trailers(cx))
.await .await
.map_err(|e| Status::from_error(&e))?; .map_err(|e| Status::from_error(Box::new(e)));
Ok(map.map(MetadataMap::from_headers)) map.map(|x| x.map(MetadataMap::from_headers))
} }
fn decode_chunk(&mut self) -> Result<Option<T>, Status> { fn decode_chunk(&mut self) -> Result<Option<T>, Status> {
@@ -232,7 +232,7 @@ impl<T> Stream for Streaming<T> {
Some(Err(e)) => { Some(Err(e)) => {
let err: crate::Error = e.into(); let err: crate::Error = e.into();
debug!("decoder inner stream error: {:?}", err); debug!("decoder inner stream error: {:?}", err);
let status = Status::from_error(&*err); let status = Status::from_error(err);
return Poll::Ready(Some(Err(status))); return Poll::Ready(Some(Err(status)));
} }
None => None, None => None,
@@ -266,7 +266,7 @@ impl<T> Stream for Streaming<T> {
Err(e) => { Err(e) => {
let err: crate::Error = e.into(); let err: crate::Error = e.into();
debug!("decoder inner trailers error: {:?}", err); debug!("decoder inner trailers error: {:?}", err);
let status = Status::from_error(&*err); let status = Status::from_error(err);
return Some(Err(status)).into(); return Some(Err(status)).into();
} }
} }
+1 -1
View File
@@ -116,7 +116,7 @@ mod tests {
let msg = Vec::from(&[0u8; 1024][..]); let msg = Vec::from(&[0u8; 1024][..]);
let messages = std::iter::repeat(Ok::<_, Status>(msg)).take(10000); let messages = std::iter::repeat_with(move || Ok::<_, Status>(msg.clone())).take(10000);
let source = futures_util::stream::iter(messages); let source = futures_util::stream::iter(messages);
let body = encode_server(encoder, source); let body = encode_server(encoder, source);
+83 -35
View File
@@ -33,7 +33,6 @@ const GRPC_STATUS_DETAILS_HEADER: &str = "grpc-status-details-bin";
/// assert_eq!(status1.code(), Code::InvalidArgument); /// assert_eq!(status1.code(), Code::InvalidArgument);
/// assert_eq!(status1.code(), status2.code()); /// assert_eq!(status1.code(), status2.code());
/// ``` /// ```
#[derive(Clone)]
pub struct Status { pub struct Status {
/// The gRPC status code, found in the `grpc-status` header. /// The gRPC status code, found in the `grpc-status` header.
code: Code, code: Code,
@@ -45,6 +44,8 @@ pub struct Status {
/// If the metadata contains any headers with names reserved either by the gRPC spec /// If the metadata contains any headers with names reserved either by the gRPC spec
/// or by `Status` fields above, they will be ignored. /// or by `Status` fields above, they will be ignored.
metadata: MetadataMap, metadata: MetadataMap,
/// Optional underlying error.
source: Option<Box<dyn Error + Send + Sync + 'static>>,
} }
/// gRPC status codes used by [`Status`]. /// gRPC status codes used by [`Status`].
@@ -162,6 +163,7 @@ impl Status {
message: message.into(), message: message.into(),
details: Bytes::new(), details: Bytes::new(),
metadata: MetadataMap::new(), metadata: MetadataMap::new(),
source: None,
} }
} }
@@ -302,38 +304,34 @@ impl Status {
} }
#[cfg_attr(not(feature = "transport"), allow(dead_code))] #[cfg_attr(not(feature = "transport"), allow(dead_code))]
pub(crate) fn from_error(err: &(dyn Error + 'static)) -> Status { pub(crate) fn from_error(err: Box<dyn Error + Send + Sync + 'static>) -> Status {
Status::try_from_error(err).unwrap_or_else(|| Status::new(Code::Unknown, err.to_string())) Status::try_from_error(err)
.unwrap_or_else(|err| Status::new(Code::Unknown, err.to_string()))
} }
pub(crate) fn try_from_error(err: &(dyn Error + 'static)) -> Option<Status> { pub(crate) fn try_from_error(
let mut cause = Some(err); err: Box<dyn Error + Send + Sync + 'static>,
) -> Result<Status, Box<dyn Error + Send + Sync + 'static>> {
while let Some(err) = cause { let err = match err.downcast::<Status>() {
if let Some(status) = err.downcast_ref::<Status>() { Ok(status) => {
return Some(Status { return Ok(*status);
code: status.code,
message: status.message.clone(),
details: status.details.clone(),
metadata: status.metadata.clone(),
});
} }
Err(err) => err,
};
#[cfg(feature = "transport")] #[cfg(feature = "transport")]
{ let err = match err.downcast::<h2::Error>() {
if let Some(h2) = err.downcast_ref::<h2::Error>() { Ok(h2) => {
return Some(Status::from_h2_error(h2)); return Ok(Status::from_h2_error(&*h2));
}
if let Some(timeout) = err.downcast_ref::<crate::transport::TimeoutExpired>() {
return Some(Status::cancelled(timeout.to_string()));
}
} }
Err(err) => err,
};
cause = err.source(); if let Some(status) = find_status_in_source_chain(&*err) {
return Ok(status);
} }
None Err(err)
} }
// FIXME: bubble this into `transport` and expose generic http2 reasons. // FIXME: bubble this into `transport` and expose generic http2 reasons.
@@ -356,7 +354,13 @@ impl Status {
_ => Code::Unknown, _ => Code::Unknown,
}; };
Status::new(code, format!("h2 protocol error: {}", err)) let mut status = Self::new(code, format!("h2 protocol error: {}", err));
let error = err
.reason()
.map(h2::Error::from)
.map(|err| Box::new(err) as Box<dyn Error + Send + Sync + 'static>);
status.source = error;
status
} }
#[cfg(feature = "transport")] #[cfg(feature = "transport")]
@@ -374,7 +378,8 @@ impl Status {
where where
E: Into<Box<dyn Error + Send + Sync>>, E: Into<Box<dyn Error + Send + Sync>>,
{ {
Status::from_error(&*err.into()) let err: Box<dyn Error + Send + Sync> = err.into();
Status::from_error(err)
} }
/// Extract a `Status` from a hyper `HeaderMap`. /// Extract a `Status` from a hyper `HeaderMap`.
@@ -410,6 +415,7 @@ impl Status {
message, message,
details, details,
metadata: MetadataMap::from_headers(other_headers), metadata: MetadataMap::from_headers(other_headers),
source: None,
}, },
Err(err) => { Err(err) => {
warn!("Error deserializing status message header: {}", err); warn!("Error deserializing status message header: {}", err);
@@ -418,6 +424,7 @@ impl Status {
message: format!("Error deserializing status message header: {}", err), message: format!("Error deserializing status message header: {}", err),
details, details,
metadata: MetadataMap::from_headers(other_headers), metadata: MetadataMap::from_headers(other_headers),
source: None,
} }
} }
} }
@@ -505,6 +512,7 @@ impl Status {
message: message.into(), message: message.into(),
details, details,
metadata, metadata,
source: None,
} }
} }
@@ -524,6 +532,32 @@ impl Status {
} }
} }
fn find_status_in_source_chain(err: &(dyn Error + 'static)) -> Option<Status> {
let mut source = Some(err);
while let Some(err) = source {
if let Some(status) = err.downcast_ref::<Status>() {
return Some(Status {
code: status.code,
message: status.message.clone(),
details: status.details.clone(),
metadata: status.metadata.clone(),
// Since `Status` is not `Clone`, any `source` on the original Status
// cannot be cloned so must remain with the original `Status`.
source: None,
});
}
#[cfg(feature = "transport")]
if let Some(timeout) = err.downcast_ref::<crate::transport::TimeoutExpired>() {
return Some(Status::cancelled(timeout.to_string()));
}
source = err.source();
}
None
}
impl fmt::Debug for Status { impl fmt::Debug for Status {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
// A manual impl to reduce the noise of frequently empty fields. // A manual impl to reduce the noise of frequently empty fields.
@@ -543,6 +577,8 @@ impl fmt::Debug for Status {
builder.field("metadata", &self.metadata); builder.field("metadata", &self.metadata);
} }
builder.field("source", &self.source);
builder.finish() builder.finish()
} }
} }
@@ -609,7 +645,11 @@ impl fmt::Display for Status {
} }
} }
impl Error for Status {} impl Error for Status {
fn source(&self) -> Option<&(dyn Error + 'static)> {
self.source.as_ref().map(|err| (&**err) as _)
}
}
/// ///
/// Take the `Status` value from `trailers` if it is available, else from `status_code`. /// Take the `Status` value from `trailers` if it is available, else from `status_code`.
@@ -775,25 +815,25 @@ mod tests {
#[test] #[test]
fn from_error_status() { fn from_error_status() {
let orig = Status::new(Code::OutOfRange, "weeaboo"); let orig = Status::new(Code::OutOfRange, "weeaboo");
let found = Status::from_error(&orig); let found = Status::from_error(Box::new(orig));
assert_eq!(orig.code(), found.code()); assert_eq!(found.code(), Code::OutOfRange);
assert_eq!(orig.message(), found.message()); assert_eq!(found.message(), "weeaboo");
} }
#[test] #[test]
fn from_error_unknown() { fn from_error_unknown() {
let orig: Error = "peek-a-boo".into(); let orig: Error = "peek-a-boo".into();
let found = Status::from_error(&*orig); let found = Status::from_error(orig);
assert_eq!(found.code(), Code::Unknown); assert_eq!(found.code(), Code::Unknown);
assert_eq!(found.message(), orig.to_string()); assert_eq!(found.message(), "peek-a-boo".to_string());
} }
#[test] #[test]
fn from_error_nested() { fn from_error_nested() {
let orig = Nested(Box::new(Status::new(Code::OutOfRange, "weeaboo"))); let orig = Nested(Box::new(Status::new(Code::OutOfRange, "weeaboo")));
let found = Status::from_error(&orig); let found = Status::from_error(Box::new(orig));
assert_eq!(found.code(), Code::OutOfRange); assert_eq!(found.code(), Code::OutOfRange);
assert_eq!(found.message(), "weeaboo"); assert_eq!(found.message(), "weeaboo");
@@ -802,10 +842,18 @@ mod tests {
#[test] #[test]
#[cfg(feature = "transport")] #[cfg(feature = "transport")]
fn from_error_h2() { fn from_error_h2() {
use std::error::Error as _;
let orig = h2::Error::from(h2::Reason::CANCEL); let orig = h2::Error::from(h2::Reason::CANCEL);
let found = Status::from_error(&orig); let found = Status::from_error(Box::new(orig));
assert_eq!(found.code(), Code::Cancelled); assert_eq!(found.code(), Code::Cancelled);
let source = found
.source()
.and_then(|err| err.downcast_ref::<h2::Error>())
.unwrap();
assert_eq!(source.reason(), Some(h2::Reason::CANCEL));
} }
#[test] #[test]
+4 -5
View File
@@ -67,15 +67,14 @@ where
let response = response.map(MaybeEmptyBody::full); let response = response.map(MaybeEmptyBody::full);
Poll::Ready(Ok(response)) Poll::Ready(Ok(response))
} }
Err(err) => { Err(err) => match Status::try_from_error(err) {
if let Some(status) = Status::try_from_error(&*err) { Ok(status) => {
let mut res = Response::new(MaybeEmptyBody::empty()); let mut res = Response::new(MaybeEmptyBody::empty());
status.add_header(res.headers_mut()).unwrap(); status.add_header(res.headers_mut()).unwrap();
Poll::Ready(Ok(res)) Poll::Ready(Ok(res))
} else {
Poll::Ready(Err(err))
} }
} Err(err) => Poll::Ready(Err(err)),
},
} }
} }
} }