use std::pin::Pin;
use std::sync::Arc;
use a2a_protocol_types::proto as apb;
use a2a_protocol_types::proto::convert::ConvertError;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use tonic::{Request, Response, Status};
use super::helpers::{server_error_to_status, validated_metadata};
use super::pb::a2a_service_server::A2aService;
use super::GrpcConfig;
use crate::handler::{RequestHandler, SendMessageResult};
type NativeStream =
Pin<Box<dyn tokio_stream::Stream<Item = Result<apb::StreamResponse, Status>> + Send + 'static>>;
#[allow(clippy::needless_pass_by_value)]
fn bad_request(err: ConvertError) -> Status {
Status::invalid_argument(err.to_string())
}
#[allow(clippy::needless_pass_by_value)]
fn bad_response(err: ConvertError) -> Status {
Status::internal(format!("response not representable in protobuf: {err}"))
}
fn send_result_to_stream(
resp: a2a_protocol_types::responses::SendMessageResponse,
) -> Result<apb::StreamResponse, ConvertError> {
let payload = match resp {
a2a_protocol_types::responses::SendMessageResponse::Task(t) => {
apb::stream_response::Payload::Task(t.try_into()?)
}
a2a_protocol_types::responses::SendMessageResponse::Message(m) => {
apb::stream_response::Payload::Message(m.try_into()?)
}
other => {
return Err(ConvertError {
field: "sendMessageResponse.payload",
reason: format!("unsupported response variant: {other:?}"),
})
}
};
Ok(apb::StreamResponse {
payload: Some(payload),
})
}
fn reader_to_native_stream(
mut reader: crate::streaming::InMemoryQueueReader,
capacity: usize,
) -> NativeStream {
use crate::streaming::EventQueueReader;
let (tx, rx) = mpsc::channel(capacity);
tokio::spawn(async move {
loop {
match reader.read().await {
Some(Ok(event)) => {
let item = apb::StreamResponse::try_from(event).map_err(bad_response);
let is_err = item.is_err();
if tx.send(item).await.is_err() || is_err {
break;
}
}
Some(Err(_)) => {
let _ = tx.send(Err(Status::internal("event queue error"))).await;
break;
}
None => break,
}
}
});
Box::pin(ReceiverStream::new(rx))
}
pub struct A2aServiceImpl {
pub(super) handler: Arc<RequestHandler>,
pub(super) config: GrpcConfig,
}
#[tonic::async_trait]
impl A2aService for A2aServiceImpl {
async fn send_message(
&self,
request: Request<apb::SendMessageRequest>,
) -> Result<Response<apb::SendMessageResponse>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::MessageSendParams =
request.into_inner().try_into().map_err(bad_request)?;
match self
.handler
.on_send_message(params, false, Some(&headers))
.await
{
Ok(SendMessageResult::Response(resp)) => {
Ok(Response::new(resp.try_into().map_err(bad_response)?))
}
Ok(SendMessageResult::Stream(_)) => Err(Status::internal(
"unexpected stream response for unary call",
)),
Err(e) => Err(server_error_to_status(&e)),
}
}
type SendStreamingMessageStream = NativeStream;
async fn send_streaming_message(
&self,
request: Request<apb::SendMessageRequest>,
) -> Result<Response<Self::SendStreamingMessageStream>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::MessageSendParams =
request.into_inner().try_into().map_err(bad_request)?;
match self
.handler
.on_send_message(params, true, Some(&headers))
.await
{
Ok(SendMessageResult::Stream(reader)) => Ok(Response::new(reader_to_native_stream(
reader,
self.config.stream_channel_capacity,
))),
Ok(SendMessageResult::Response(resp)) => {
let payload = send_result_to_stream(resp).map_err(bad_response)?;
let stream = Box::pin(tokio_stream::once(Ok(payload)));
Ok(Response::new(stream as NativeStream))
}
Err(e) => Err(server_error_to_status(&e)),
}
}
async fn get_task(
&self,
request: Request<apb::GetTaskRequest>,
) -> Result<Response<apb::Task>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::TaskQueryParams =
request.into_inner().try_into().map_err(bad_request)?;
match self.handler.on_get_task(params, Some(&headers)).await {
Ok(task) => Ok(Response::new(task.try_into().map_err(bad_response)?)),
Err(e) => Err(server_error_to_status(&e)),
}
}
async fn list_tasks(
&self,
request: Request<apb::ListTasksRequest>,
) -> Result<Response<apb::ListTasksResponse>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::ListTasksParams =
request.into_inner().try_into().map_err(bad_request)?;
match self.handler.on_list_tasks(params, Some(&headers)).await {
Ok(resp) => Ok(Response::new(resp.try_into().map_err(bad_response)?)),
Err(e) => Err(server_error_to_status(&e)),
}
}
async fn cancel_task(
&self,
request: Request<apb::CancelTaskRequest>,
) -> Result<Response<apb::Task>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::CancelTaskParams =
request.into_inner().try_into().map_err(bad_request)?;
match self.handler.on_cancel_task(params, Some(&headers)).await {
Ok(task) => Ok(Response::new(task.try_into().map_err(bad_response)?)),
Err(e) => Err(server_error_to_status(&e)),
}
}
type SubscribeToTaskStream = NativeStream;
async fn subscribe_to_task(
&self,
request: Request<apb::SubscribeToTaskRequest>,
) -> Result<Response<Self::SubscribeToTaskStream>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::TaskIdParams = request.into_inner().into();
match self.handler.on_resubscribe(params, Some(&headers)).await {
Ok(reader) => Ok(Response::new(reader_to_native_stream(
reader,
self.config.stream_channel_capacity,
))),
Err(e) => Err(server_error_to_status(&e)),
}
}
async fn create_task_push_notification_config(
&self,
request: Request<apb::TaskPushNotificationConfig>,
) -> Result<Response<apb::TaskPushNotificationConfig>, Status> {
let headers = validated_metadata(request.metadata())?;
let config: a2a_protocol_types::push::TaskPushNotificationConfig =
request.into_inner().into();
match self
.handler
.on_set_push_config(config, Some(&headers))
.await
{
Ok(cfg) => Ok(Response::new(cfg.into())),
Err(e) => Err(server_error_to_status(&e)),
}
}
async fn get_task_push_notification_config(
&self,
request: Request<apb::GetTaskPushNotificationConfigRequest>,
) -> Result<Response<apb::TaskPushNotificationConfig>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::GetPushConfigParams = request.into_inner().into();
match self
.handler
.on_get_push_config(params, Some(&headers))
.await
{
Ok(cfg) => Ok(Response::new(cfg.into())),
Err(e) => Err(server_error_to_status(&e)),
}
}
async fn list_task_push_notification_configs(
&self,
request: Request<apb::ListTaskPushNotificationConfigsRequest>,
) -> Result<Response<apb::ListTaskPushNotificationConfigsResponse>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::ListPushConfigsParams =
request.into_inner().try_into().map_err(bad_request)?;
match self
.handler
.on_list_push_configs(¶ms.task_id, params.tenant.as_deref(), Some(&headers))
.await
{
Ok(configs) => Ok(Response::new(
apb::ListTaskPushNotificationConfigsResponse {
configs: configs.into_iter().map(Into::into).collect(),
next_page_token: String::new(),
},
)),
Err(e) => Err(server_error_to_status(&e)),
}
}
async fn delete_task_push_notification_config(
&self,
request: Request<apb::DeleteTaskPushNotificationConfigRequest>,
) -> Result<Response<()>, Status> {
let headers = validated_metadata(request.metadata())?;
let params: a2a_protocol_types::params::DeletePushConfigParams =
request.into_inner().into();
match self
.handler
.on_delete_push_config(params, Some(&headers))
.await
{
Ok(()) => Ok(Response::new(())),
Err(e) => Err(server_error_to_status(&e)),
}
}
async fn get_extended_agent_card(
&self,
request: Request<apb::GetExtendedAgentCardRequest>,
) -> Result<Response<apb::AgentCard>, Status> {
let headers = validated_metadata(request.metadata())?;
match self
.handler
.on_get_extended_agent_card(Some(&headers))
.await
{
Ok(card) => Ok(Response::new(card.try_into().map_err(bad_response)?)),
Err(e) => Err(server_error_to_status(&e)),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use a2a_protocol_types::message::{Message, MessageId, MessageRole, Part, PartContent};
use a2a_protocol_types::responses::SendMessageResponse;
use a2a_protocol_types::task::{ContextId, TaskId};
#[test]
fn bad_request_maps_to_invalid_argument() {
let status = bad_request(ConvertError {
field: "message.role",
reason: "unknown Role number 9".into(),
});
assert_eq!(status.code(), tonic::Code::InvalidArgument);
assert!(status.message().contains("message.role"));
}
#[test]
fn bad_response_maps_to_internal() {
let status = bad_response(ConvertError {
field: "task.metadata",
reason: "boom".into(),
});
assert_eq!(status.code(), tonic::Code::Internal);
}
#[test]
fn send_result_message_wraps_into_stream_payload() {
let resp = SendMessageResponse::Message(Message {
id: MessageId("m".into()),
role: MessageRole::Agent,
parts: vec![Part {
content: PartContent::Text("hi".into()),
metadata: None,
filename: None,
media_type: None,
}],
task_id: Some(TaskId("t".into())),
context_id: Some(ContextId("c".into())),
reference_task_ids: None,
extensions: None,
metadata: None,
});
let stream = send_result_to_stream(resp).unwrap();
assert!(matches!(
stream.payload,
Some(apb::stream_response::Payload::Message(_))
));
}
}