use pingora_core::protocols::http::HttpTask;
use pingora_proxy::Session;
use praxis_core::grpc::{GrpcStatusCode, encode_grpc_message};
use praxis_filter::GrpcErrorMapping;
use tracing::debug;
pub(crate) async fn send_trailers_only(
session: &mut Session,
mapping: &GrpcErrorMapping,
http_status: u16,
message: &str,
) {
let grpc_status = GrpcStatusCode::from_http_status(http_status);
let Some(header) = build_header(session, mapping, grpc_status, message) else {
return;
};
debug!(
http_status,
grpc_status = grpc_status.as_u32(),
"sending gRPC trailers-only response"
);
write_header_only(session, header).await;
}
pub(crate) async fn send_grpc_rejection(session: &mut Session, rejection: &praxis_filter::Rejection) {
let mut header = match pingora_http::ResponseHeader::build(rejection.status, Some(rejection.headers.len())) {
Ok(header) => header,
Err(error) => {
debug!(%error, "could not build a gRPC rejection header");
return;
},
};
for (name, value) in &rejection.headers {
if session.req_header().version == http::Version::HTTP_2 && name.eq_ignore_ascii_case("content-length") {
continue;
}
if let Err(error) = header.append_header(name.clone(), value.clone()) {
debug!(%error, name, "dropping unencodable gRPC rejection header");
}
}
debug!("sending gRPC rejection as a trailers-only response");
write_header_only(session, header).await;
}
async fn write_header_only(session: &mut Session, header: pingora_http::ResponseHeader) {
if let Err(error) = session
.as_downstream_mut()
.response_duplex_vec(vec![HttpTask::Header(Box::new(header), true)])
.await
{
debug!(%error, "failed to write gRPC trailers-only response");
}
}
fn build_header(
session: &Session,
mapping: &GrpcErrorMapping,
grpc_status: GrpcStatusCode,
message: &str,
) -> Option<pingora_http::ResponseHeader> {
let mut header = pingora_http::ResponseHeader::build(200, Some(4))
.inspect_err(|error| debug!(%error, "could not build a gRPC error response header"))
.ok()?;
let _insert = header.insert_header(http::header::CONTENT_TYPE, mapping.content_type().clone());
let _insert = header.insert_header("grpc-status", grpc_status.as_header_value());
if mapping.include_message() && !message.is_empty() {
let encoded = encode_grpc_message(message);
match http::HeaderValue::from_str(&encoded) {
Ok(value) => {
let _insert = header.insert_header("grpc-message", value);
},
Err(error) => debug!(%error, "dropping unencodable grpc-message"),
}
}
if session.req_header().version != http::Version::HTTP_2 {
let _insert = header.insert_header(http::header::CONTENT_LENGTH, "0");
}
Some(header)
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(clippy::unwrap_used, clippy::expect_used, reason = "tests")]
mod tests {
use praxis_core::grpc::GrpcKind;
use super::*;
fn make_test_mapping_with_message() -> GrpcErrorMapping {
GrpcErrorMapping::new(GrpcKind::Grpc, true)
}
fn make_test_mapping_without_message() -> GrpcErrorMapping {
GrpcErrorMapping::new(GrpcKind::GrpcProto, false)
}
#[test]
fn grpc_error_mapping_include_message() {
let mapping_with = make_test_mapping_with_message();
assert!(mapping_with.include_message());
let mapping_without = make_test_mapping_without_message();
assert!(!mapping_without.include_message());
}
#[test]
fn grpc_status_as_header_value() -> Result<(), Box<dyn std::error::Error>> {
let codes = vec![
GrpcStatusCode::Ok,
GrpcStatusCode::Cancelled,
GrpcStatusCode::Unknown,
GrpcStatusCode::InvalidArgument,
GrpcStatusCode::DeadlineExceeded,
GrpcStatusCode::NotFound,
GrpcStatusCode::AlreadyExists,
GrpcStatusCode::PermissionDenied,
GrpcStatusCode::ResourceExhausted,
GrpcStatusCode::FailedPrecondition,
GrpcStatusCode::Aborted,
GrpcStatusCode::OutOfRange,
GrpcStatusCode::Unimplemented,
GrpcStatusCode::Internal,
GrpcStatusCode::Unavailable,
GrpcStatusCode::DataLoss,
GrpcStatusCode::Unauthenticated,
];
for code in codes {
let header_val = code.as_header_value();
let val_str = header_val.to_str()?;
let parsed = val_str.parse::<u32>()?;
assert_eq!(parsed, code.as_u32(), "Invalid gRPC status: {val_str}");
}
Ok(())
}
#[test]
fn edge_case_http_status_codes() {
for status in [0, 100, 301, 999] {
assert_eq!(
GrpcStatusCode::from_http_status(status),
GrpcStatusCode::Unknown,
"unmapped HTTP status {status} should fall back to UNKNOWN without panicking"
);
}
}
#[test]
fn message_encoding_edge_cases() {
let long_message = format!("Error: {}", "x".repeat(1000));
let encoded = encode_grpc_message(&long_message);
assert!(
encoded.len() >= long_message.len(),
"encoding a very long message must not shrink it"
);
let special_only = "\n\r\t";
let encoded = encode_grpc_message(special_only);
assert!(
!encoded.is_empty(),
"a message of only special characters must still encode"
);
assert!(encoded.starts_with('%'), "control characters must be percent-encoded");
let with_null = "Error\0Details";
let encoded = encode_grpc_message(with_null);
assert!(encoded.contains("%00"), "a null byte must be percent-encoded to %00");
let already_encoded = "Error%20message";
let encoded = encode_grpc_message(already_encoded);
assert!(
encoded.contains("%25"),
"an already percent-encoded message has its own percent sign encoded"
);
}
#[test]
fn content_type_variations() {
let kinds = vec![
(GrpcKind::Grpc, "application/grpc"),
(GrpcKind::GrpcProto, "application/grpc+proto"),
(GrpcKind::GrpcJson, "application/grpc+json"),
];
for (kind, expected_ct) in kinds {
let mapping = GrpcErrorMapping::new(kind, true);
assert_eq!(mapping.content_type().to_str().unwrap(), expected_ct);
}
}
}