use cratestack_core::{CoolError, rpc::RpcErrorBody};
use reqwest::StatusCode;
use serde::de::DeserializeOwned;
use crate::codec::HttpClientCodec;
use crate::error::ClientError;
use crate::runtime::wire::RuntimeResponseWire;
#[derive(Debug, Clone)]
pub struct RpcRemoteError {
pub status: StatusCode,
pub body: RpcErrorBody,
}
impl std::fmt::Display for RpcRemoteError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"RPC call failed with code {} (status {}): {}",
self.body.code,
self.status.as_u16(),
self.body.message
)
}
}
impl std::error::Error for RpcRemoteError {}
#[derive(Debug, thiserror::Error)]
pub enum RpcClientError {
#[error("transport error: {0}")]
Transport(#[from] reqwest::Error),
#[error("codec error: {0}")]
Codec(#[from] CoolError),
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error("bad input: {0}")]
BadInput(String),
#[error("{0}")]
Remote(RpcRemoteError),
}
pub type RpcStream<O> = tokio::sync::mpsc::Receiver<Result<O, RpcClientError>>;
pub(crate) fn http_status_for_rpc_code(code: &str) -> StatusCode {
match code {
"invalid_argument" => StatusCode::BAD_REQUEST,
"unauthenticated" => StatusCode::UNAUTHORIZED,
"permission_denied" => StatusCode::FORBIDDEN,
"not_found" => StatusCode::NOT_FOUND,
"conflict" => StatusCode::CONFLICT,
"failed_precondition" => StatusCode::PRECONDITION_FAILED,
_ => StatusCode::INTERNAL_SERVER_ERROR,
}
}
pub(crate) fn client_error_to_rpc(error: ClientError) -> RpcClientError {
match error {
ClientError::Transport(error) => RpcClientError::Transport(error),
ClientError::Codec(error) => RpcClientError::Codec(error),
ClientError::InvalidResponse(message) => RpcClientError::InvalidResponse(message),
ClientError::BadInput(message) => RpcClientError::BadInput(message),
ClientError::State(message) => RpcClientError::InvalidResponse(message),
ClientError::Remote {
status,
error,
message,
} => {
let body = error
.map(cratestack_core::rpc::RpcErrorBody::from_cool_response)
.unwrap_or_else(|| RpcErrorBody {
code: "internal".to_owned(),
message,
details: None,
});
RpcClientError::Remote(RpcRemoteError { status, body })
}
}
}
pub(crate) fn decode_rpc_unary_response<C, Output>(
codec: &C,
response: &RuntimeResponseWire,
) -> Result<Output, RpcClientError>
where
C: HttpClientCodec,
Output: DeserializeOwned,
{
let content_type = response
.headers
.iter()
.find(|header| header.name.eq_ignore_ascii_case("content-type"))
.map(|header| header.value.as_str())
.ok_or_else(|| {
RpcClientError::InvalidResponse("response is missing Content-Type header".to_owned())
})?;
if (200..=299).contains(&response.status_code) {
codec
.decode_response::<Output>(content_type, &response.body)
.map_err(RpcClientError::Codec)
} else {
let body = codec
.decode_response::<RpcErrorBody>(content_type, &response.body)
.unwrap_or_else(|_| RpcErrorBody {
code: "internal".to_owned(),
message: format!(
"unexpected RPC error body for status {}",
response.status_code
),
details: None,
});
Err(RpcClientError::Remote(RpcRemoteError {
status: StatusCode::from_u16(response.status_code)
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR),
body,
}))
}
}