use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use super::native::A2aServiceImpl;
use super::{A2aServiceServer, GrpcConfig};
use crate::handler::RequestHandler;
pub struct GrpcDispatcher {
handler: Arc<RequestHandler>,
config: GrpcConfig,
keepalive: Option<(Duration, Duration)>,
max_connection_age: Option<Duration>,
}
impl GrpcDispatcher {
#[must_use]
pub const fn new(handler: Arc<RequestHandler>, config: GrpcConfig) -> Self {
Self {
handler,
config,
keepalive: None,
max_connection_age: None,
}
}
#[must_use]
pub const fn with_http2_keepalive(mut self, interval: Duration, timeout: Duration) -> Self {
self.keepalive = Some((interval, timeout));
self
}
#[must_use]
pub const fn with_max_connection_age(mut self, age: Duration) -> Self {
self.max_connection_age = Some(age);
self
}
pub async fn serve(self, addr: impl tokio::net::ToSocketAddrs) -> std::io::Result<()> {
let addr = super::helpers::resolve_addr(addr).await?;
trace_info!(
addr = %addr,
"A2A gRPC server listening"
);
let router = self.build_router();
router.serve(addr).await.map_err(std::io::Error::other)
}
pub async fn serve_with_addr(
self,
addr: impl tokio::net::ToSocketAddrs,
) -> std::io::Result<SocketAddr> {
let listener = tokio::net::TcpListener::bind(addr).await?;
self.serve_with_listener(listener)
}
pub fn serve_with_listener(
self,
listener: tokio::net::TcpListener,
) -> std::io::Result<SocketAddr> {
let local_addr = listener.local_addr()?;
let incoming = tokio_stream::wrappers::TcpListenerStream::new(listener);
trace_info!(
%local_addr,
"A2A gRPC server listening"
);
let router = self.build_router();
tokio::spawn(async move {
let _ = router.serve_with_incoming(incoming).await;
});
Ok(local_addr)
}
#[must_use]
pub fn into_service(&self) -> A2aServiceServer<A2aServiceImpl> {
let inner = A2aServiceImpl {
handler: Arc::clone(&self.handler),
config: self.config.clone(),
};
A2aServiceServer::new(inner)
.max_decoding_message_size(self.config.max_message_size)
.max_encoding_message_size(self.config.max_message_size)
}
fn build_router(&self) -> tonic::transport::server::Router {
let mut server = tonic::transport::Server::builder()
.concurrency_limit_per_connection(self.config.concurrency_limit);
if let Some((interval, timeout)) = self.keepalive {
server = server
.http2_keepalive_interval(Some(interval))
.http2_keepalive_timeout(Some(timeout));
}
if let Some(age) = self.max_connection_age {
server = server.max_connection_age(age);
}
server.add_service(self.into_service())
}
}
impl std::fmt::Debug for GrpcDispatcher {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GrpcDispatcher")
.field("handler", &"RequestHandler { .. }")
.field("config", &self.config)
.field("keepalive", &self.keepalive)
.field("max_connection_age", &self.max_connection_age)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn connection_knobs_are_off_by_default_and_carry_what_they_are_given() {
use crate::agent_executor;
use crate::RequestHandlerBuilder;
use std::sync::Arc;
struct DummyExec;
agent_executor!(DummyExec, |_ctx, _queue| async { Ok(()) });
let handler = Arc::new(RequestHandlerBuilder::new(DummyExec).build().unwrap());
let default = GrpcDispatcher::new(Arc::clone(&handler), GrpcConfig::default());
assert!(
default.keepalive.is_none(),
"HTTP/2 keepalive must be opt-in: enabling it by default would start \
pinging the clients of every deployment that upgrades"
);
assert!(
default.max_connection_age.is_none(),
"and so must connection ageing, which forces reconnects"
);
let tuned = GrpcDispatcher::new(handler, GrpcConfig::default())
.with_http2_keepalive(Duration::from_secs(30), Duration::from_secs(10))
.with_max_connection_age(Duration::from_secs(600));
assert_eq!(
tuned.keepalive,
Some((Duration::from_secs(30), Duration::from_secs(10)))
);
assert_eq!(tuned.max_connection_age, Some(Duration::from_secs(600)));
let _router = tuned.build_router();
}
#[test]
fn grpc_dispatcher_debug_does_not_panic() {
use crate::agent_executor;
use crate::RequestHandlerBuilder;
use std::sync::Arc;
struct DummyExec;
agent_executor!(DummyExec, |_ctx, _queue| async { Ok(()) });
let handler = Arc::new(RequestHandlerBuilder::new(DummyExec).build().unwrap());
let dispatcher = GrpcDispatcher::new(handler, GrpcConfig::default());
let _ = format!("{dispatcher:?}");
}
}