use parking_lot::{Mutex, RwLock};
use std::collections::{BTreeMap, HashMap};
use std::future::Future;
#[cfg(not(target_arch = "wasm32"))]
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, Instant};
use crate::bytes::Bytes;
#[cfg(not(target_arch = "wasm32"))]
use crate::bytes::BytesMut;
use crate::cx::{Cx, cap};
#[cfg(not(target_arch = "wasm32"))]
use crate::http::h1::server::HostPolicy;
#[cfg(not(target_arch = "wasm32"))]
use crate::http::h1::types::{Request as HttpRequest, Response as HttpResponse};
#[cfg(not(target_arch = "wasm32"))]
use crate::http::h2::listener::{Http2Listener, Http2ListenerConfig};
#[cfg(not(target_arch = "wasm32"))]
use crate::http::h2::settings::Settings;
#[cfg(not(target_arch = "wasm32"))]
use crate::runtime::RuntimeHandle;
#[cfg(not(target_arch = "wasm32"))]
use crate::server::shutdown::ShutdownStats;
#[cfg(not(target_arch = "wasm32"))]
use base64::Engine as _;
use super::client::CompressionEncoding;
pub use super::codec::RequestBodyMeter;
use super::codec::{Codec, FramedCodec};
use super::reflection::ReflectionService;
use super::service::{NamedService, ServiceHandler};
use super::status::{GrpcError, Status, TransportErrorKind};
use super::streaming::{Metadata, Request, Response};
fn wall_clock_instant_now() -> Instant {
Instant::now()
}
async fn poll_with_current_cx<F: Future>(cx: Cx, future: F) -> F::Output {
let mut future = std::pin::pin!(future);
std::future::poll_fn(|task_cx| {
let _guard = Cx::set_current(Some(cx.clone()));
future.as_mut().poll(task_cx)
})
.await
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct InclusiveDeadlineElapsed;
async fn poll_before_inclusive_deadline<F: Future>(
future: F,
deadline: crate::types::Time,
now: fn() -> crate::types::Time,
) -> Result<F::Output, InclusiveDeadlineElapsed> {
let mut future = std::pin::pin!(future);
std::future::poll_fn(|task_cx| {
if now() >= deadline {
std::task::Poll::Ready(Err(InclusiveDeadlineElapsed))
} else {
future.as_mut().poll(task_cx).map(Ok)
}
})
.await
}
async fn invoke_and_poll_before_inclusive_deadline<H, A, F>(
handler: H,
argument: A,
deadline: crate::types::Time,
now: fn() -> crate::types::Time,
) -> Result<F::Output, InclusiveDeadlineElapsed>
where
H: FnOnce(A) -> F,
F: Future,
{
if now() >= deadline {
return Err(InclusiveDeadlineElapsed);
}
let future = handler(argument);
poll_before_inclusive_deadline(future, deadline, now).await
}
#[derive(Debug, Clone)]
struct StreamState {
last_activity: Instant,
registered_at: Instant,
}
#[derive(Debug)]
pub struct ConnectionState {
active_streams: Mutex<HashMap<u32, StreamState>>,
}
impl ConnectionState {
pub fn new() -> Self {
Self {
active_streams: Mutex::new(HashMap::new()),
}
}
pub fn add_stream(&self, stream_id: u32, max_concurrent: u32) -> Result<Instant, String> {
let mut active_streams = self.active_streams.lock();
if active_streams.len() >= max_concurrent as usize {
return Err(format!(
"connection exceeds max_concurrent_streams: {} >= {}",
active_streams.len(),
max_concurrent
));
}
let now = wall_clock_instant_now();
active_streams.insert(
stream_id,
StreamState {
last_activity: now,
registered_at: now,
},
);
Ok(now)
}
pub fn update_stream_activity(&self, stream_id: u32) {
let mut active_streams = self.active_streams.lock();
if let Some(stream) = active_streams.get_mut(&stream_id) {
stream.last_activity = wall_clock_instant_now();
}
}
pub fn remove_stream(&self, stream_id: u32) {
let mut active_streams = self.active_streams.lock();
active_streams.remove(&stream_id);
}
pub fn cleanup_idle_streams(&self, idle_timeout: Duration) -> Vec<u32> {
let now = wall_clock_instant_now();
let mut removed = Vec::new();
let mut active_streams = self.active_streams.lock();
active_streams.retain(|&stream_id, stream| {
let idle_duration = now.duration_since(stream.last_activity);
if idle_duration > idle_timeout {
removed.push(stream_id);
false
} else {
true
}
});
removed
}
pub fn active_stream_count(&self) -> usize {
let active_streams = self.active_streams.lock();
active_streams.len()
}
pub fn remove_stream_if_owned(&self, stream_id: u32, registered_at: Instant) {
let mut active_streams = self.active_streams.lock();
if let Some(stream_state) = active_streams.get(&stream_id) {
if stream_state.registered_at == registered_at {
active_streams.remove(&stream_id);
}
}
}
}
#[derive(Debug)]
pub struct ConnectionRegistry {
connections: RwLock<HashMap<String, ConnectionState>>,
}
impl ConnectionRegistry {
pub fn new() -> Self {
Self {
connections: RwLock::new(HashMap::new()),
}
}
pub fn add_connection(&self, connection_id: String) {
let mut connections = self.connections.write();
connections.insert(connection_id, ConnectionState::new());
}
pub fn remove_connection(&self, connection_id: &str) {
let mut connections = self.connections.write();
connections.remove(connection_id);
}
pub fn enforce_stream_limits(
&self,
connection_id: &str,
stream_id: u32,
max_concurrent: u32,
idle_timeout: Option<Duration>,
) -> Result<Instant, String> {
let connections = self.connections.read();
let connection = connections
.get(connection_id)
.ok_or_else(|| format!("connection not registered: {}", connection_id))?;
if let Some(timeout) = idle_timeout {
connection.cleanup_idle_streams(timeout);
}
connection.add_stream(stream_id, max_concurrent)
}
pub fn update_stream_activity(&self, connection_id: &str, stream_id: u32) {
let connections = self.connections.read();
if let Some(connection) = connections.get(connection_id) {
connection.update_stream_activity(stream_id);
}
}
pub fn remove_stream(&self, connection_id: &str, stream_id: u32) {
let connections = self.connections.read();
if let Some(connection) = connections.get(connection_id) {
connection.remove_stream(stream_id);
}
}
pub fn remove_stream_if_owned(
&self,
connection_id: &str,
stream_id: u32,
registered_at: Instant,
) {
let connections = self.connections.read();
if let Some(connection) = connections.get(connection_id) {
connection.remove_stream_if_owned(stream_id, registered_at);
}
}
pub fn get_stats(&self) -> (usize, usize) {
let connections = self.connections.read();
let connection_count = connections.len();
let total_streams: usize = connections
.values()
.map(|conn| conn.active_stream_count())
.sum();
(connection_count, total_streams)
}
}
struct StreamRegistrationGuard {
registry: Arc<ConnectionRegistry>,
connection_id: String,
stream_id: u32,
registered_at: Instant,
}
impl Drop for StreamRegistrationGuard {
fn drop(&mut self) {
self.registry.remove_stream_if_owned(
&self.connection_id,
self.stream_id,
self.registered_at,
);
}
}
#[derive(Debug, Clone)]
pub struct ServerConfig {
pub max_recv_message_size: usize,
pub max_send_message_size: usize,
pub max_request_body_bytes: Option<usize>,
pub initial_connection_window_size: u32,
pub initial_stream_window_size: u32,
pub max_concurrent_streams: u32,
pub keepalive_interval_ms: Option<u64>,
pub keepalive_timeout_ms: Option<u64>,
pub default_timeout: Option<Duration>,
pub max_request_deadline: Option<Duration>,
pub send_compression: Option<CompressionEncoding>,
pub accept_compression: Vec<CompressionEncoding>,
pub max_metadata_size: usize,
pub stream_idle_timeout: Option<Duration>,
}
pub const DEFAULT_MAX_METADATA_SIZE: usize = 8 * 1024;
#[must_use]
pub fn metadata_byte_size(metadata: &super::streaming::Metadata) -> usize {
let mut total = 0usize;
for (key, value) in metadata.iter() {
let value_len = match value {
super::streaming::MetadataValue::Ascii(s) => s.len(),
super::streaming::MetadataValue::Binary(b) => b.len(),
};
total = total.saturating_add(key.len()).saturating_add(value_len);
}
total
}
fn metadata_key_uses_grpc_prefix(key: &str) -> bool {
key.get(..5)
.is_some_and(|prefix| prefix.eq_ignore_ascii_case("grpc-"))
}
fn grpc_request_header_is_allowed(key: &str) -> bool {
key.eq_ignore_ascii_case("grpc-timeout")
|| key.eq_ignore_ascii_case("grpc-encoding")
|| key.eq_ignore_ascii_case("grpc-accept-encoding")
|| key.eq_ignore_ascii_case("grpc-message-type")
}
fn matches_media_type_prefix(value: &str, prefix: &str) -> bool {
value.starts_with(prefix)
&& matches!(value.as_bytes().get(prefix.len()), None | Some(b'+' | b';'))
}
fn grpc_content_type_is_allowed(value: &str) -> bool {
matches_media_type_prefix(value.trim(), "application/grpc")
}
fn grpc_te_header_is_allowed(value: &str) -> bool {
value.trim().eq_ignore_ascii_case("trailers")
}
fn is_valid_header_name_rfc7230(name: &str) -> bool {
if name.is_empty() {
return false;
}
for byte in name.bytes() {
match byte {
b'A'..=b'Z' | b'a'..=b'z' => {}
b'0'..=b'9' => {}
b'!' | b'#' | b'$' | b'%' | b'&' | b'\'' | b'*' | b'+' | b'-' | b'.' | b'^' | b'_'
| b'`' | b'|' | b'~' => {}
_ => return false,
}
}
true
}
fn is_valid_header_value_rfc7230(value: &str) -> bool {
let bytes = value.as_bytes();
if value.contains('\r') || value.contains('\n') {
return false;
}
for &byte in bytes {
match byte {
0x21..=0x7E => {}
b' ' | b'\t' => {}
_ => return false,
}
}
true
}
const MAX_HEADER_NAME_LEN: usize = 256; const MAX_HEADER_VALUE_LEN: usize = 8192;
fn validate_inbound_metadata(metadata: &super::streaming::Metadata) -> Result<(), Status> {
for (key, value) in metadata.iter() {
if !is_valid_header_name_rfc7230(key) {
return Err(Status::invalid_argument(format!(
"metadata key '{key}' contains invalid characters (RFC 7230 violation)"
)));
}
if key.len() > MAX_HEADER_NAME_LEN {
return Err(Status::invalid_argument(format!(
"metadata key '{key}' exceeds maximum length ({} > {})",
key.len(),
MAX_HEADER_NAME_LEN
)));
}
match value {
super::streaming::MetadataValue::Ascii(text) => {
if !is_valid_header_value_rfc7230(text) {
return Err(Status::invalid_argument(format!(
"metadata value for '{key}' contains disallowed CRLF or invalid characters (RFC 7230 violation)"
)));
}
if text.len() > MAX_HEADER_VALUE_LEN {
return Err(Status::invalid_argument(format!(
"metadata value for '{key}' exceeds maximum length ({} > {})",
text.len(),
MAX_HEADER_VALUE_LEN
)));
}
}
super::streaming::MetadataValue::Binary(bytes) => {
if bytes.len() > MAX_HEADER_VALUE_LEN {
return Err(Status::invalid_argument(format!(
"binary metadata value for '{key}' exceeds maximum length ({} > {})",
bytes.len(),
MAX_HEADER_VALUE_LEN
)));
}
}
}
if metadata_key_uses_grpc_prefix(key) && !grpc_request_header_is_allowed(key) {
return Err(Status::invalid_argument(format!(
"client metadata key uses reserved grpc-* prefix: {key}"
)));
}
if let super::streaming::MetadataValue::Ascii(text) = value {
if super::streaming::sanitize_metadata_ascii_value(text).as_ref() != text {
return Err(Status::invalid_argument(format!(
"metadata value for {key} contains disallowed control or non-ASCII bytes"
)));
}
}
if key.eq_ignore_ascii_case("content-type") {
match value {
super::streaming::MetadataValue::Ascii(text)
if !grpc_content_type_is_allowed(text) =>
{
return Err(Status::invalid_argument(format!(
"content-type must be application/grpc(+proto|+json), got {text}"
)));
}
super::streaming::MetadataValue::Binary(_) => {
return Err(Status::invalid_argument(
"content-type must be an ASCII gRPC media type",
));
}
super::streaming::MetadataValue::Ascii(_) => {}
}
} else if key.eq_ignore_ascii_case("te") {
match value {
super::streaming::MetadataValue::Ascii(text)
if !grpc_te_header_is_allowed(text) =>
{
return Err(Status::invalid_argument(format!(
"te must be trailers for gRPC over HTTP/2, got {text}"
)));
}
super::streaming::MetadataValue::Binary(_) => {
return Err(Status::invalid_argument(
"te must be an ASCII trailers header",
));
}
super::streaming::MetadataValue::Ascii(_) => {}
}
}
}
Ok(())
}
pub fn enforce_metadata_size_limit(
metadata: &super::streaming::Metadata,
limit: usize,
) -> Result<(), Status> {
validate_inbound_metadata(metadata)?;
if limit == 0 {
return Ok(());
}
let actual = metadata_byte_size(metadata);
if actual > limit {
return Err(Status::resource_exhausted(format!(
"metadata exceeds max_metadata_size: {actual} bytes > {limit} bytes \
(gRPC equivalent of HTTP 431 Request Header Fields Too Large; \
see ServerConfig::max_metadata_size)"
)));
}
Ok(())
}
#[cfg(not(target_arch = "wasm32"))]
fn grpc_request_trailer_key_is_reserved(key: &str) -> bool {
metadata_key_uses_grpc_prefix(key)
|| [
"age",
"authorization",
"cache-control",
"connection",
"content-encoding",
"content-length",
"content-range",
"content-type",
"cookie",
"date",
"expect",
"expires",
"host",
"keep-alive",
"max-forwards",
"pragma",
"proxy-authenticate",
"proxy-authorization",
"proxy-connection",
"range",
"retry-after",
"set-cookie",
"te",
"trailer",
"transfer-encoding",
"upgrade",
"vary",
"warning",
"www-authenticate",
]
.iter()
.any(|reserved| key.eq_ignore_ascii_case(reserved))
}
#[cfg(not(target_arch = "wasm32"))]
fn insert_http2_metadata_entry(
metadata: &mut Metadata,
key: &str,
value: &str,
) -> Result<(), Status> {
let binary = key
.get(key.len().saturating_sub(4)..)
.is_some_and(|suffix| suffix.eq_ignore_ascii_case("-bin"));
let inserted = if binary {
let decoded = base64::engine::general_purpose::STANDARD
.decode(value)
.or_else(|_| base64::engine::general_purpose::STANDARD_NO_PAD.decode(value))
.map_err(|_| {
Status::invalid_argument(format!(
"binary metadata value for '{key}' is not valid base64"
))
})?;
metadata.insert_bin(key, Bytes::from(decoded))
} else {
metadata.insert(key, value)
};
if inserted {
Ok(())
} else {
Err(Status::invalid_argument(format!(
"invalid gRPC metadata entry '{key}'"
)))
}
}
#[cfg(not(target_arch = "wasm32"))]
fn enforce_http2_metadata_blocks(
headers: &Metadata,
trailers: &Metadata,
limit: usize,
) -> Result<(), Status> {
validate_inbound_metadata(headers)?;
validate_inbound_metadata(trailers)?;
if limit == 0 {
return Ok(());
}
let actual = metadata_byte_size(headers).saturating_add(metadata_byte_size(trailers));
if actual > limit {
return Err(Status::resource_exhausted(format!(
"combined request headers and trailers exceed max_metadata_size: \
{actual} bytes > {limit} bytes"
)));
}
Ok(())
}
impl RequestBodyMeter {
#[must_use]
pub fn from_config(config: &ServerConfig) -> Self {
Self::new(config.max_request_body_bytes)
}
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
max_recv_message_size: 4 * 1024 * 1024, max_send_message_size: 4 * 1024 * 1024, max_request_body_bytes: None,
initial_connection_window_size: 1024 * 1024,
initial_stream_window_size: 1024 * 1024,
max_concurrent_streams: 100,
keepalive_interval_ms: None,
keepalive_timeout_ms: None,
default_timeout: None,
max_request_deadline: None,
send_compression: None,
accept_compression: vec![CompressionEncoding::Identity],
max_metadata_size: DEFAULT_MAX_METADATA_SIZE,
stream_idle_timeout: Some(Duration::from_secs(60)),
}
}
}
#[derive(Default)]
pub struct ServerBuilder {
config: ServerConfig,
services: BTreeMap<String, Arc<dyn ServiceHandler>>,
reflection: Option<ReflectionService>,
interceptors: Vec<Arc<dyn Interceptor>>,
}
impl std::fmt::Debug for ServerBuilder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ServerBuilder")
.field("config", &self.config)
.field("services", &format!("[{} services]", self.services.len()))
.field("reflection_enabled", &self.reflection.is_some())
.finish()
}
}
impl ServerBuilder {
#[must_use]
pub fn new() -> Self {
Self {
config: ServerConfig::default(),
services: BTreeMap::new(),
reflection: None,
interceptors: Vec::new(),
}
}
#[must_use]
pub fn interceptor<I>(mut self, interceptor: I) -> Self
where
I: Interceptor + 'static,
{
self.interceptors.push(Arc::new(interceptor));
self
}
#[must_use]
pub fn interceptor_arc(mut self, interceptor: Arc<dyn Interceptor>) -> Self {
self.interceptors.push(interceptor);
self
}
#[must_use]
pub fn max_recv_message_size(mut self, size: usize) -> Self {
self.config.max_recv_message_size = size;
self
}
#[must_use]
pub fn max_metadata_size(mut self, size: usize) -> Self {
self.config.max_metadata_size = size;
self
}
#[must_use]
pub fn stream_idle_timeout(mut self, timeout: Option<Duration>) -> Self {
self.config.stream_idle_timeout = timeout;
self
}
#[must_use]
pub fn max_send_message_size(mut self, size: usize) -> Self {
self.config.max_send_message_size = size;
self
}
#[must_use]
pub fn max_request_body_bytes(mut self, size: usize) -> Self {
self.config.max_request_body_bytes = Some(size);
self
}
#[must_use]
pub fn initial_connection_window_size(mut self, size: u32) -> Self {
self.config.initial_connection_window_size = size;
self
}
#[must_use]
pub fn initial_stream_window_size(mut self, size: u32) -> Self {
self.config.initial_stream_window_size = size;
self
}
#[must_use]
pub fn max_concurrent_streams(mut self, max: u32) -> Self {
self.config.max_concurrent_streams = max;
self
}
#[must_use]
pub fn keepalive_interval(mut self, ms: u64) -> Self {
self.config.keepalive_interval_ms = Some(ms);
self
}
#[must_use]
pub fn keepalive_timeout(mut self, ms: u64) -> Self {
self.config.keepalive_timeout_ms = Some(ms);
self
}
#[must_use]
pub fn default_timeout(mut self, timeout: Duration) -> Self {
self.config.default_timeout = Some(timeout);
self
}
#[must_use]
pub fn max_request_deadline(mut self, max: Duration) -> Self {
self.config.max_request_deadline = Some(max);
self
}
#[must_use]
pub fn send_compression(mut self, encoding: CompressionEncoding) -> Self {
self.config.send_compression = Some(encoding);
self
}
#[must_use]
pub fn accept_compression(mut self, encoding: CompressionEncoding) -> Self {
self.config.accept_compression.push(encoding);
self
}
#[must_use]
pub fn accept_compressions(
mut self,
encodings: impl IntoIterator<Item = CompressionEncoding>,
) -> Self {
self.config.accept_compression.clear();
self.config.accept_compression.extend(encodings);
self
}
#[must_use]
pub fn add_service<S>(mut self, service: S) -> Self
where
S: NamedService + ServiceHandler + 'static,
{
let service_name = S::NAME.to_string();
let service: Arc<dyn ServiceHandler> = Arc::new(service);
if let Some(reflection) = self.reflection.as_ref()
&& service_name != ReflectionService::NAME
{
reflection.register_handler(service.as_ref());
}
self.services.insert(service_name, service);
self
}
#[must_use]
pub fn enable_reflection_with_auth<F>(mut self, auth_callback: F) -> Self
where
F: Fn(&Cx, &str) -> Result<(), Status> + Send + Sync + 'static,
{
let reflection = self
.reflection
.take()
.unwrap_or_default()
.with_auth(auth_callback);
for service in self.services.values() {
if service.descriptor().full_name() != ReflectionService::NAME {
reflection.register_handler(service.as_ref());
}
}
self.services.insert(
ReflectionService::NAME.to_string(),
Arc::new(reflection.clone()),
);
self.reflection = Some(reflection);
self
}
#[deprecated(
since = "0.3.3",
note = "Use enable_reflection_with_auth() to install production reflection auth explicitly"
)]
#[must_use]
pub fn enable_reflection(mut self) -> Self {
let reflection = self.reflection.take().unwrap_or_default(); for service in self.services.values() {
if service.descriptor().full_name() != ReflectionService::NAME {
reflection.register_handler(service.as_ref());
}
}
self.services.insert(
ReflectionService::NAME.to_string(),
Arc::new(reflection.clone()),
);
self.reflection = Some(reflection);
self
}
#[must_use]
pub fn build(self) -> Server {
Server {
config: self.config,
services: self.services,
interceptors: self.interceptors,
connection_registry: Arc::new(ConnectionRegistry::new()),
}
}
}
pub struct Server {
config: ServerConfig,
services: BTreeMap<String, Arc<dyn ServiceHandler>>,
interceptors: Vec<Arc<dyn Interceptor>>,
connection_registry: Arc<ConnectionRegistry>,
}
#[cfg(not(target_arch = "wasm32"))]
#[derive(Debug)]
pub struct GrpcTransportRequest {
path: String,
request: Request<Bytes>,
trailing_metadata: Metadata,
}
#[cfg(not(target_arch = "wasm32"))]
impl GrpcTransportRequest {
#[must_use]
pub fn path(&self) -> &str {
&self.path
}
#[must_use]
pub fn request(&self) -> &Request<Bytes> {
&self.request
}
#[must_use]
pub fn trailing_metadata(&self) -> &Metadata {
&self.trailing_metadata
}
#[must_use]
pub fn into_parts(self) -> (String, Request<Bytes>, Metadata) {
(self.path, self.request, self.trailing_metadata)
}
}
#[cfg(not(target_arch = "wasm32"))]
type GrpcHttp2Future = Pin<Box<dyn Future<Output = HttpResponse> + Send>>;
impl std::fmt::Debug for Server {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Server")
.field("config", &self.config)
.field("services", &format!("[{} services]", self.services.len()))
.finish()
}
}
impl Server {
#[must_use]
pub fn builder() -> ServerBuilder {
ServerBuilder::new()
}
#[must_use]
pub fn config(&self) -> &ServerConfig {
&self.config
}
#[must_use]
pub fn framed_codec<C: Codec>(&self, inner: C) -> FramedCodec<C> {
FramedCodec::with_message_size_limits(
inner,
self.config.max_send_message_size,
self.config.max_recv_message_size,
)
.with_request_body_limit(self.config.max_request_body_bytes)
}
#[cfg(not(target_arch = "wasm32"))]
#[must_use]
pub fn http2_listener_config(&self, host_policy: HostPolicy) -> Http2ListenerConfig {
let mut settings = Settings::server();
settings.initial_window_size = self.config.initial_stream_window_size;
settings.max_concurrent_streams = self.config.max_concurrent_streams;
settings.max_header_list_size = if self.config.max_metadata_size == 0 {
u32::MAX
} else {
u32::try_from(self.config.max_metadata_size).unwrap_or(u32::MAX)
};
let max_body_size = self
.config
.max_recv_message_size
.saturating_add(super::codec::MESSAGE_HEADER_SIZE);
Http2ListenerConfig::default()
.settings(settings)
.initial_connection_window_size(self.config.initial_connection_window_size)
.max_body_size(max_body_size)
.host_policy(host_policy)
.stream_idle_timeout(self.config.stream_idle_timeout)
}
#[cfg(not(target_arch = "wasm32"))]
fn validate_http2_transport_config(&self) -> io::Result<()> {
if self.config.initial_stream_window_size > 0x7fff_ffff {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"gRPC initial stream window exceeds the HTTP/2 31-bit maximum",
));
}
if !(65_535..=0x7fff_ffff).contains(&self.config.initial_connection_window_size) {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"gRPC initial connection window must be within 65535..=2^31-1",
));
}
Ok(())
}
#[cfg(not(target_arch = "wasm32"))]
pub async fn bind_http2<A, F, Fut>(
self: &Arc<Self>,
addr: A,
host_policy: HostPolicy,
handler: F,
) -> io::Result<Http2Listener<impl Fn(HttpRequest) -> GrpcHttp2Future + Send + Sync + 'static>>
where
A: std::net::ToSocketAddrs + Send + 'static,
F: Fn(GrpcTransportRequest) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Response<Bytes>, Status>> + Send + 'static,
{
self.validate_http2_transport_config()?;
let config = self.http2_listener_config(host_policy);
let server = Arc::clone(self);
let handler = Arc::new(handler);
let transport_handler = move |request: HttpRequest| -> GrpcHttp2Future {
let server = Arc::clone(&server);
let handler = Arc::clone(&handler);
Box::pin(async move { server.dispatch_http2_unary(request, handler).await })
};
Http2Listener::bind_with_config(addr, transport_handler, config).await
}
#[cfg(not(target_arch = "wasm32"))]
pub async fn bind_registered_http2<A>(
self: &Arc<Self>,
addr: A,
host_policy: HostPolicy,
) -> io::Result<Http2Listener<impl Fn(HttpRequest) -> GrpcHttp2Future + Send + Sync + 'static>>
where
A: std::net::ToSocketAddrs + Send + 'static,
{
if self.services.is_empty() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"cannot bind registered gRPC routing without a service",
));
}
self.validate_http2_transport_config()?;
let config = self.http2_listener_config(host_policy);
let server = Arc::clone(self);
let transport_handler = move |request: HttpRequest| -> GrpcHttp2Future {
let server = Arc::clone(&server);
Box::pin(async move { server.dispatch_http2_registered_unary(request).await })
};
Http2Listener::bind_with_config(addr, transport_handler, config).await
}
#[cfg(not(target_arch = "wasm32"))]
pub async fn serve_http2<A>(
self: &Arc<Self>,
runtime: &RuntimeHandle,
addr: A,
host_policy: HostPolicy,
) -> io::Result<ShutdownStats>
where
A: std::net::ToSocketAddrs + Send + 'static,
{
self.bind_registered_http2(addr, host_policy)
.await?
.run(runtime)
.await
}
#[cfg(not(target_arch = "wasm32"))]
async fn dispatch_http2_unary<F, Fut>(
&self,
request: HttpRequest,
handler: Arc<F>,
) -> HttpResponse
where
F: Fn(GrpcTransportRequest) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Response<Bytes>, Status>> + Send + 'static,
{
let (path, request, trailing_metadata) = match self.decode_http2_unary_request(request) {
Ok(decoded) => decoded,
Err(status) => return Self::http2_status_response(&status),
};
let result = self
.dispatch_unary(request, move |request| {
handler(GrpcTransportRequest {
path,
request,
trailing_metadata,
})
})
.await;
match result {
Ok(response) => match self.encode_http2_unary_response(&response) {
Ok(response) => response,
Err(status) => Self::http2_status_response(&status),
},
Err(status) => Self::http2_status_response(&status),
}
}
#[cfg(not(target_arch = "wasm32"))]
async fn dispatch_http2_registered_unary(&self, request: HttpRequest) -> HttpResponse {
let (path, request, trailing_metadata) = match self.decode_http2_unary_request(request) {
Ok(decoded) => decoded,
Err(status) => return Self::http2_status_response(&status),
};
let Some(cx) = Cx::current() else {
return Self::http2_status_response(&Status::internal(
"registered gRPC HTTP/2 dispatch requires a runtime Cx",
));
};
let result = self
.dispatch_registered_unary_with_trailers(&cx, &path, request, trailing_metadata)
.await;
match result {
Ok(response) => match self.encode_http2_unary_response(&response) {
Ok(response) => response,
Err(status) => Self::http2_status_response(&status),
},
Err(status) => Self::http2_status_response(&status),
}
}
#[cfg(not(target_arch = "wasm32"))]
fn decode_http2_unary_request(
&self,
request: HttpRequest,
) -> Result<(String, Request<Bytes>, Metadata), Status> {
if request.method != crate::http::h1::types::Method::Post {
return Err(Status::invalid_argument(
"gRPC over HTTP/2 requires the POST method",
));
}
let content_type = request
.content_type()
.ok_or_else(|| Status::invalid_argument("missing gRPC content-type"))?;
if !grpc_content_type_is_allowed(content_type) {
return Err(Status::invalid_argument(format!(
"content-type must be application/grpc(+proto|+json), got {content_type}"
)));
}
let mut metadata = Metadata::new();
metadata.reserve(request.headers.len());
for (key, value) in &request.headers {
insert_http2_metadata_entry(&mut metadata, key, value)?;
}
let mut trailing_metadata = Metadata::new();
trailing_metadata.reserve(request.trailers.len());
let mut trailer_keys = std::collections::BTreeSet::new();
for (key, value) in &request.trailers {
let normalized_key = key.to_ascii_lowercase();
if grpc_request_trailer_key_is_reserved(key) {
return Err(Status::invalid_argument(format!(
"gRPC request trailer uses reserved transport key '{key}'"
)));
}
if !trailer_keys.insert(normalized_key) || metadata.get(key).is_some() {
return Err(Status::invalid_argument(format!(
"duplicate gRPC request trailer metadata key '{key}'"
)));
}
insert_http2_metadata_entry(&mut trailing_metadata, key, value)?;
}
enforce_http2_metadata_blocks(
&metadata,
&trailing_metadata,
self.config.max_metadata_size,
)?;
let grpc_encoding = request.header_value("grpc-encoding");
let encoding = match grpc_encoding {
Some(value) => CompressionEncoding::from_header_value(value).ok_or_else(|| {
Status::unimplemented(format!("unsupported grpc-encoding: {value}"))
})?,
None => CompressionEncoding::Identity,
};
if !self.config.accept_compression.contains(&encoding) {
return Err(Status::unimplemented(format!(
"grpc-encoding is not accepted by this server: {}",
grpc_encoding.unwrap_or("identity")
)));
}
let mut codec = self.framed_codec(super::codec::IdentityCodec);
if encoding != CompressionEncoding::Identity {
let decompressor = encoding.frame_decompressor().ok_or_else(|| {
Status::unimplemented(format!(
"grpc-encoding support is not compiled in: {}",
grpc_encoding.unwrap_or("identity")
))
})?;
codec = codec.with_frame_hooks(None, Some(decompressor));
}
let mut body = BytesMut::from(request.body.as_slice());
let message = codec
.decode_message_with_encoding(&mut body, grpc_encoding)
.map_err(GrpcError::into_status)?
.ok_or_else(|| Status::invalid_argument("incomplete gRPC message frame"))?;
if !body.is_empty() {
match codec.decode_message_with_encoding(&mut body, grpc_encoding) {
Ok(Some(_)) => {
return Err(Status::invalid_argument(
"unary gRPC request contains more than one message",
));
}
Ok(None) => {
return Err(Status::invalid_argument(
"unary gRPC request has a truncated trailing frame",
));
}
Err(error) => return Err(error.into_status()),
}
}
Ok((
request.uri,
Request::with_metadata(message, metadata),
trailing_metadata,
))
}
#[cfg(not(target_arch = "wasm32"))]
fn encode_http2_unary_response(
&self,
response: &Response<Bytes>,
) -> Result<HttpResponse, Status> {
let mut codec = self.framed_codec(super::codec::IdentityCodec);
let mut grpc_encoding = None;
if let Some(encoding) = self.config.send_compression {
if encoding != CompressionEncoding::Identity {
let compressor = encoding.frame_compressor().ok_or_else(|| {
Status::unimplemented("configured response compression is not compiled in")
})?;
codec = codec.with_frame_hooks(Some(compressor), None);
grpc_encoding = Some(match encoding {
CompressionEncoding::Identity => "identity",
CompressionEncoding::Gzip => "gzip",
});
}
}
let mut body = BytesMut::new();
codec
.encode_message(response.get_ref(), &mut body)
.map_err(GrpcError::into_status)?;
let mut http = HttpResponse::new(200, "OK", body.to_vec())
.with_header("content-type", "application/grpc");
if let Some(encoding) = grpc_encoding {
http.headers
.push(("grpc-encoding".to_owned(), encoding.to_owned()));
}
for (key, value) in response.metadata().iter() {
if key.eq_ignore_ascii_case("content-type")
|| key.eq_ignore_ascii_case("grpc-status")
|| key.eq_ignore_ascii_case("grpc-message")
{
return Err(Status::internal(format!(
"response metadata uses transport-reserved key '{key}'"
)));
}
let value = match value {
super::streaming::MetadataValue::Ascii(value) => value.clone(),
super::streaming::MetadataValue::Binary(value) => {
base64::engine::general_purpose::STANDARD_NO_PAD.encode(value)
}
};
http.headers.push((key.to_owned(), value));
}
http.trailers
.push(("grpc-status".to_owned(), "0".to_owned()));
Ok(http)
}
#[cfg(not(target_arch = "wasm32"))]
fn http2_status_response(status: &Status) -> HttpResponse {
let mut response = HttpResponse::new(200, "OK", Vec::new())
.with_header("content-type", "application/grpc");
if !status.message().is_empty() {
response.trailers.push((
"grpc-message".to_owned(),
super::status::percent_encode_grpc_message(status.message()),
));
}
if let Some(details) = status.details() {
response.trailers.push((
"grpc-status-details-bin".to_owned(),
base64::engine::general_purpose::STANDARD_NO_PAD.encode(details),
));
}
response
.trailers
.push(("grpc-status".to_owned(), status.code().as_i32().to_string()));
response
}
#[must_use]
pub fn services(&self) -> &BTreeMap<String, Arc<dyn ServiceHandler>> {
&self.services
}
#[must_use]
pub fn connection_registry(&self) -> &Arc<ConnectionRegistry> {
&self.connection_registry
}
pub fn register_connection(&self, connection_id: String) {
self.connection_registry.add_connection(connection_id);
}
pub fn unregister_connection(&self, connection_id: &str) {
self.connection_registry.remove_connection(connection_id);
}
fn clear_auth_context_from_request(request: &mut Request<Bytes>) {
let _ = request
.extensions_mut()
.remove_typed::<super::interceptor::AuthContext>();
}
#[must_use]
pub fn interceptors(&self) -> &[Arc<dyn Interceptor>] {
&self.interceptors
}
pub async fn dispatch_unary<H, F>(
&self,
mut request: Request<Bytes>,
handler: H,
) -> Result<Response<Bytes>, Status>
where
H: FnOnce(Request<Bytes>) -> F,
F: Future<Output = Result<Response<Bytes>, Status>>,
{
enforce_metadata_size_limit(request.metadata(), self.config.max_metadata_size)?;
RequestBodyMeter::from_config(&self.config)
.record_message_bytes(request.get_ref().len())?;
for (index, interceptor) in self.interceptors.iter().enumerate() {
if let Err(mut status) = interceptor.intercept_request(&mut request) {
for cleanup in self.interceptors[..=index].iter().rev() {
if let Err(replacement) =
cleanup.intercept_error_with_request(&request, &mut status)
{
status = replacement;
}
}
Self::clear_auth_context_from_request(&mut request);
return Err(status);
}
}
let call_context = CallContext::from_metadata_at_with_max_deadline(
request.metadata().clone(),
self.config.default_timeout,
self.config.max_request_deadline,
None, wall_clock_instant_now(),
);
let mut request_snapshot = request.snapshot(Bytes::new());
let response_result = if call_context.deadline().is_some() {
let time_now = crate::time::wall_now();
let now = wall_clock_instant_now();
let Some(remaining_duration) = call_context.remaining_at(now) else {
Self::clear_auth_context_from_request(&mut request_snapshot);
return Err(Status::deadline_exceeded(
"Request deadline already expired",
));
};
let runtime_deadline = time_now + remaining_duration;
let base_budget =
Cx::current().map_or(crate::types::Budget::INFINITE, |ambient| ambient.budget());
let source = if grpc_timeout_from_metadata(request.metadata()).is_some() {
crate::web::request_region::RequestBudgetSource::HeaderClamped
} else {
crate::web::request_region::RequestBudgetSource::ServerConfig
};
let budget = base_budget.tightened_by_timeout(time_now, remaining_duration);
let region =
crate::web::request_region::ServerRequestRegion::mint("h2-grpc", budget, time_now);
let handler_future = invoke_and_poll_before_inclusive_deadline(
handler,
request,
runtime_deadline,
crate::time::wall_now,
);
match region {
Some(region) => {
let scoped = region.instrumented(source, handler_future);
match crate::time::timeout_at(runtime_deadline, scoped).await {
Ok(Ok(result)) => {
region.finish(if result.is_ok() { "ok" } else { "err" });
result
}
Ok(Err(_)) | Err(_) => {
region.cancel_timeout("grpc request deadline exceeded");
region.finish("deadline_exceeded");
Self::clear_auth_context_from_request(&mut request_snapshot);
return Err(Status::deadline_exceeded("Request deadline exceeded"));
}
}
}
None => {
match crate::time::timeout_at(runtime_deadline, handler_future).await {
Ok(Ok(result)) => result,
Ok(Err(_)) | Err(_) => {
Self::clear_auth_context_from_request(&mut request_snapshot);
return Err(Status::deadline_exceeded("Request deadline exceeded"));
}
}
}
}
} else {
handler(request).await
};
let mut response = match response_result {
Ok(response) => response,
Err(mut status) => {
for interceptor in self.interceptors.iter().rev() {
if let Err(replacement) =
interceptor.intercept_error_with_request(&request_snapshot, &mut status)
{
status = replacement;
}
}
Self::clear_auth_context_from_request(&mut request_snapshot);
return Err(status);
}
};
for interceptor in self.interceptors.iter().rev() {
if let Err(mut status) =
interceptor.intercept_response_with_request(&request_snapshot, &mut response)
{
for cleanup in self.interceptors.iter().rev() {
if let Err(replacement) =
cleanup.intercept_error_with_request(&request_snapshot, &mut status)
{
status = replacement;
}
}
Self::clear_auth_context_from_request(&mut request_snapshot);
return Err(status);
}
}
Ok(response)
}
pub async fn dispatch_unary_with_stream_enforcement<H, F>(
&self,
connection_id: String,
stream_id: u32,
request: Request<Bytes>,
handler: H,
) -> Result<Response<Bytes>, Status>
where
H: FnOnce(Request<Bytes>) -> F,
F: Future<Output = Result<Response<Bytes>, Status>>,
{
let registered_at = match self.connection_registry.enforce_stream_limits(
&connection_id,
stream_id,
self.config.max_concurrent_streams,
self.config.stream_idle_timeout,
) {
Ok(timestamp) => timestamp,
Err(limit_error) => {
return Err(Status::resource_exhausted(format!(
"stream limit enforcement failed: {}",
limit_error
)));
}
};
let _stream_guard = StreamRegistrationGuard {
registry: Arc::clone(&self.connection_registry),
connection_id: connection_id.clone(),
stream_id,
registered_at,
};
self.dispatch_unary(request, handler).await
}
pub fn update_stream_activity(&self, connection_id: &str, stream_id: u32) {
self.connection_registry
.update_stream_activity(connection_id, stream_id);
}
pub fn get_connection_stats(&self) -> (usize, usize) {
self.connection_registry.get_stats()
}
fn resolve_registered_unary(&self, path: &str) -> Result<Arc<dyn ServiceHandler>, Status> {
let Some(route) = path.strip_prefix('/') else {
return Err(Status::unimplemented(format!(
"unknown gRPC method path '{path}'"
)));
};
let mut segments = route.split('/');
let (Some(service_name), Some(method_name), None) =
(segments.next(), segments.next(), segments.next())
else {
return Err(Status::unimplemented(format!(
"unknown gRPC method path '{path}'"
)));
};
if service_name.is_empty() || method_name.is_empty() {
return Err(Status::unimplemented(format!(
"unknown gRPC method path '{path}'"
)));
}
let service = self.services.get(service_name).ok_or_else(|| {
Status::unimplemented(format!("gRPC service '{service_name}' is not registered"))
})?;
let method = service
.descriptor()
.methods
.iter()
.find(|method| method.path == path && method.name == method_name)
.ok_or_else(|| {
Status::unimplemented(format!("gRPC method '{path}' is not registered"))
})?;
if !method.is_unary() {
return Err(Status::unimplemented(format!(
"gRPC method '{path}' is streaming and cannot use unary dispatch"
)));
}
Ok(Arc::clone(service))
}
pub async fn dispatch_registered_unary(
&self,
cx: &Cx,
path: &str,
request: Request<Bytes>,
) -> Result<Response<Bytes>, Status> {
self.dispatch_registered_unary_with_trailers(cx, path, request, Metadata::new())
.await
}
pub async fn dispatch_registered_unary_with_trailers(
&self,
cx: &Cx,
path: &str,
request: Request<Bytes>,
trailing_metadata: Metadata,
) -> Result<Response<Bytes>, Status> {
if cx.checkpoint().is_err() {
let status = match cx.cancel_reason().map(|reason| reason.kind) {
Some(crate::types::CancelKind::Timeout | crate::types::CancelKind::Deadline) => {
Status::deadline_exceeded(
"registered gRPC request deadline elapsed before dispatch",
)
}
Some(
crate::types::CancelKind::PollQuota | crate::types::CancelKind::CostBudget,
) => Status::resource_exhausted(
"registered gRPC request budget was exhausted before dispatch",
),
_ => Status::cancelled("registered gRPC request was cancelled before dispatch"),
};
return Err(status);
}
let service = self.resolve_registered_unary(path)?;
let path = path.to_owned();
let base_cx = cx.clone();
let dispatch = self.dispatch_unary(request, move |request| async move {
let call_cx = Cx::current().unwrap_or_else(|| base_cx.clone());
service
.call_unary(&call_cx, &path, request, trailing_metadata)
.await
});
poll_with_current_cx(cx.clone(), dispatch).await
}
#[must_use]
pub fn get_service(&self, name: &str) -> Option<&Arc<dyn ServiceHandler>> {
self.services.get(name)
}
pub fn service_names(&self) -> Vec<&str> {
self.services.keys().map(String::as_str).collect()
}
#[allow(clippy::unused_async)]
pub async fn serve(self, addr: &str) -> Result<(), GrpcError> {
if self.services.is_empty() {
return Err(GrpcError::protocol(
"cannot serve gRPC server without registered services",
));
}
let listener = std::net::TcpListener::bind(addr).map_err(|error| {
GrpcError::transport_kind(
TransportErrorKind::from_io_error_kind(error.kind()),
format!("bind failed: {error}"),
)
})?;
listener.set_nonblocking(true).map_err(|error| {
GrpcError::transport_kind(
TransportErrorKind::from_io_error_kind(error.kind()),
format!("nonblocking setup failed: {error}"),
)
})?;
Ok(())
}
}
#[must_use]
pub fn parse_grpc_timeout(header: &str) -> Option<Duration> {
if header.is_empty() {
return None;
}
if !header.is_ascii() {
return None;
}
let (digits, unit) = header.split_at(header.len() - 1);
if digits.is_empty() || digits.len() > 8 || !digits.bytes().all(|byte| byte.is_ascii_digit()) {
return None;
}
let value: u64 = digits.parse().ok()?;
match unit {
"H" => Some(Duration::from_secs(value.checked_mul(3600)?)),
"M" => Some(Duration::from_secs(value.checked_mul(60)?)),
"S" => Some(Duration::from_secs(value)),
"m" => Some(Duration::from_millis(value)),
"u" => Some(Duration::from_micros(value)),
"n" => Some(Duration::from_nanos(value)),
_ => None,
}
}
fn grpc_timeout_from_metadata(metadata: &Metadata) -> Option<Duration> {
match metadata.get("grpc-timeout") {
Some(super::streaming::MetadataValue::Ascii(value)) => parse_grpc_timeout(value),
Some(super::streaming::MetadataValue::Binary(_)) | None => None,
}
}
#[must_use]
pub fn format_grpc_timeout(duration: Duration) -> String {
const MAX_VALUE: u128 = 99_999_999;
let ns = duration.as_nanos();
if ns == 0 {
return "0n".to_string();
}
let secs = u128::from(duration.as_secs());
if duration.subsec_nanos() == 0 {
let hours = secs / 3600;
if hours <= MAX_VALUE && secs % 3600 == 0 {
return format!("{hours}H");
}
let mins = secs / 60;
if mins <= MAX_VALUE && secs % 60 == 0 {
return format!("{mins}M");
}
if secs <= MAX_VALUE {
return format!("{secs}S");
}
}
let ms = duration.as_millis();
if ms <= MAX_VALUE && ns.is_multiple_of(1_000_000) {
return format!("{ms}m");
}
let us = duration.as_micros();
if us <= MAX_VALUE && ns.is_multiple_of(1_000) {
return format!("{us}u");
}
if ns <= MAX_VALUE {
return format!("{ns}n");
}
if us <= MAX_VALUE {
return format!("{us}u");
}
if ms <= MAX_VALUE {
return format!("{ms}m");
}
if secs <= MAX_VALUE {
return format!("{secs}S");
}
let mins = secs / 60;
if mins <= MAX_VALUE {
return format!("{mins}M");
}
let hours = (mins / 60).min(MAX_VALUE);
format!("{hours}H")
}
#[derive(Debug)]
pub struct CallContext {
metadata: Metadata,
deadline: Option<Instant>,
peer_addr: Option<String>,
time_getter: fn() -> Instant,
}
impl CallContext {
#[must_use]
pub fn new() -> Self {
Self {
metadata: Metadata::new(),
deadline: None,
peer_addr: None,
time_getter: wall_clock_instant_now,
}
}
#[must_use]
pub fn from_metadata(
metadata: Metadata,
default_timeout: Option<Duration>,
peer_addr: Option<String>,
) -> Self {
Self::from_metadata_with_time_getter(
metadata,
default_timeout,
peer_addr,
wall_clock_instant_now,
)
}
#[must_use]
pub fn from_metadata_with_time_getter(
metadata: Metadata,
default_timeout: Option<Duration>,
peer_addr: Option<String>,
time_getter: fn() -> Instant,
) -> Self {
Self::from_metadata_at(metadata, default_timeout, peer_addr, time_getter())
.with_time_getter(time_getter)
}
#[must_use]
pub fn from_metadata_at(
metadata: Metadata,
default_timeout: Option<Duration>,
peer_addr: Option<String>,
now: Instant,
) -> Self {
Self::from_metadata_at_with_max_deadline(metadata, default_timeout, None, peer_addr, now)
}
#[must_use]
pub fn from_metadata_at_with_max_deadline(
metadata: Metadata,
default_timeout: Option<Duration>,
max_request_deadline: Option<Duration>,
peer_addr: Option<String>,
now: Instant,
) -> Self {
let peer_timeout = grpc_timeout_from_metadata(&metadata);
let timeout = peer_timeout
.map(|peer| max_request_deadline.map_or(peer, |cap| peer.min(cap)))
.or(default_timeout);
let deadline = timeout.map(|t| now.checked_add(t).unwrap_or(now));
Self {
metadata,
deadline,
peer_addr,
time_getter: wall_clock_instant_now,
}
}
#[must_use]
pub fn with_deadline(deadline: Instant) -> Self {
Self {
metadata: Metadata::new(),
deadline: Some(deadline),
peer_addr: None,
time_getter: wall_clock_instant_now,
}
}
#[must_use]
pub const fn with_time_getter(mut self, time_getter: fn() -> Instant) -> Self {
self.time_getter = time_getter;
self
}
#[must_use]
pub const fn time_getter(&self) -> fn() -> Instant {
self.time_getter
}
#[must_use]
pub fn metadata(&self) -> &Metadata {
&self.metadata
}
#[must_use]
pub fn deadline(&self) -> Option<Instant> {
self.deadline
}
#[must_use]
pub fn peer_addr(&self) -> Option<&str> {
self.peer_addr.as_deref()
}
#[must_use]
pub fn remaining(&self) -> Option<Duration> {
self.remaining_at((self.time_getter)())
}
#[must_use]
pub fn remaining_at(&self, now: Instant) -> Option<Duration> {
self.deadline.and_then(|deadline| {
deadline
.checked_duration_since(now)
.filter(|remaining| !remaining.is_zero())
})
}
#[must_use]
pub fn timeout_header_value(&self) -> Option<String> {
self.timeout_header_value_at((self.time_getter)())
}
#[must_use]
pub fn timeout_header_value_at(&self, now: Instant) -> Option<String> {
self.deadline
.map(|deadline| format_grpc_timeout(deadline.saturating_duration_since(now)))
}
pub fn propagate_timeout_to(&self, metadata: &mut Metadata) -> bool {
self.propagate_timeout_to_at(metadata, (self.time_getter)())
}
pub fn propagate_timeout_to_at(&self, metadata: &mut Metadata, now: Instant) -> bool {
let Some(parent_remaining) = self
.deadline
.map(|deadline| deadline.saturating_duration_since(now))
else {
return false;
};
let effective = match metadata.get("grpc-timeout") {
Some(super::streaming::MetadataValue::Ascii(existing)) => parse_grpc_timeout(existing)
.map_or(parent_remaining, |child| child.min(parent_remaining)),
Some(super::streaming::MetadataValue::Binary(_)) | None => parent_remaining,
};
let _ = metadata.insert_or_replace("grpc-timeout", format_grpc_timeout(effective));
true
}
#[must_use]
pub fn is_expired(&self) -> bool {
self.is_expired_at((self.time_getter)())
}
#[must_use]
pub fn is_expired_at(&self, now: Instant) -> bool {
self.deadline.is_some_and(|deadline| now >= deadline)
}
#[must_use]
pub fn with_cx<'a>(&'a self, cx: &'a Cx) -> CallContextWithCx<'a> {
CallContextWithCx { call: self, cx }
}
}
impl Default for CallContext {
fn default() -> Self {
Self::new()
}
}
pub struct CallContextWithCx<'a> {
call: &'a CallContext,
cx: &'a Cx,
}
impl CallContextWithCx<'_> {
#[must_use]
pub fn call(&self) -> &CallContext {
self.call
}
#[must_use]
pub fn metadata(&self) -> &Metadata {
self.call.metadata()
}
#[must_use]
pub fn deadline(&self) -> Option<std::time::Instant> {
self.call.deadline()
}
#[must_use]
pub fn peer_addr(&self) -> Option<&str> {
self.call.peer_addr()
}
#[must_use]
pub fn is_expired(&self) -> bool {
self.call.is_expired()
}
#[must_use]
pub fn remaining(&self) -> Option<Duration> {
self.call.remaining()
}
#[must_use]
pub fn timeout_header_value(&self) -> Option<String> {
self.call.timeout_header_value()
}
pub fn propagate_timeout_to(&self, metadata: &mut Metadata) -> bool {
self.call.propagate_timeout_to(metadata)
}
#[must_use]
pub fn cx(&self) -> &Cx {
self.cx
}
#[must_use]
pub fn cx_narrow<Caps>(&self) -> Cx<Caps>
where
Caps: cap::SubsetOf<cap::All>,
{
self.cx.restrict::<Caps>()
}
#[must_use]
pub fn cx_readonly(&self) -> Cx<cap::None> {
self.cx.restrict::<cap::None>()
}
}
pub trait Interceptor: Send + Sync {
fn intercept_request(&self, request: &mut Request<Bytes>) -> Result<(), Status>;
fn intercept_response(&self, response: &mut Response<Bytes>) -> Result<(), Status>;
fn intercept_response_with_request(
&self,
request: &Request<Bytes>,
response: &mut Response<Bytes>,
) -> Result<(), Status> {
let _ = request;
self.intercept_response(response)
}
fn intercept_error_with_request(
&self,
request: &Request<Bytes>,
status: &mut Status,
) -> Result<(), Status> {
let _ = (request, status);
Ok(())
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct NoopInterceptor;
impl Interceptor for NoopInterceptor {
fn intercept_request(&self, _request: &mut Request<Bytes>) -> Result<(), Status> {
Ok(())
}
fn intercept_response(&self, _response: &mut Response<Bytes>) -> Result<(), Status> {
Ok(())
}
}
#[derive(Debug)]
pub struct AuthInterceptor<F> {
validator: F,
}
impl<F> AuthInterceptor<F>
where
F: Fn(&Metadata) -> Result<(), Status> + Send + Sync,
{
#[must_use]
pub fn new(validator: F) -> Self {
Self { validator }
}
}
impl<F> Interceptor for AuthInterceptor<F>
where
F: Fn(&Metadata) -> Result<(), Status> + Send + Sync,
{
fn intercept_request(&self, request: &mut Request<Bytes>) -> Result<(), Status> {
(self.validator)(request.metadata())
}
fn intercept_response(&self, _response: &mut Response<Bytes>) -> Result<(), Status> {
Ok(())
}
}
pub type UnaryHandler<Req, Resp> =
Box<dyn Fn(Request<Req>) -> UnaryFuture<Resp> + Send + Sync + 'static>;
pub type UnaryFuture<Resp> =
Pin<Box<dyn Future<Output = Result<Response<Resp>, Status>> + Send + 'static>>;
pub fn ok<T>(message: T) -> Result<Response<T>, Status> {
Ok(Response::new(message))
}
pub fn err<T>(status: Status) -> Result<Response<T>, Status> {
Err(status)
}
#[cfg(test)]
include!("server_tests.rs");