use std::collections::HashMap;
use std::net::SocketAddr;
use tonic::Status;
use crate::error::ServerError;
pub(super) fn extract_metadata(metadata: &tonic::metadata::MetadataMap) -> HashMap<String, String> {
let mut map = HashMap::new();
for kv in metadata.iter() {
if let tonic::metadata::KeyAndValueRef::Ascii(key, value) = kv {
if let Ok(v) = value.to_str() {
map.insert(key.as_str().to_owned(), v.to_owned());
}
}
}
map
}
#[allow(clippy::result_large_err)]
pub(super) fn validated_metadata(
metadata: &tonic::metadata::MetadataMap,
) -> Result<HashMap<String, String>, Status> {
let headers = extract_metadata(metadata);
if let Some(v) = headers.get("a2a-version") {
let v = v.trim();
if !v.is_empty() {
let major = v.split('.').next().and_then(|s| s.parse::<u32>().ok());
if major != Some(1) {
return Err(server_error_to_status(&ServerError::Protocol(
a2a_protocol_types::error::A2aError::version_not_supported(format!(
"unsupported A2A version: {v}; this server supports 1.x"
)),
)));
}
}
}
Ok(headers)
}
pub(super) fn server_error_to_status(err: &ServerError) -> Status {
use a2a_protocol_types::ErrorCode;
if let ServerError::Overloaded(msg) = err {
return Status::new(tonic::Code::ResourceExhausted, msg.clone());
}
let a2a_err = err.to_a2a_error();
let code = match a2a_err.code {
ErrorCode::TaskNotFound => tonic::Code::NotFound,
ErrorCode::TaskNotCancelable
| ErrorCode::ExtendedAgentCardNotConfigured
| ErrorCode::ExtensionSupportRequired => tonic::Code::FailedPrecondition,
ErrorCode::ContentTypeNotSupported
| ErrorCode::InvalidParams
| ErrorCode::InvalidRequest
| ErrorCode::ParseError => tonic::Code::InvalidArgument,
ErrorCode::MethodNotFound
| ErrorCode::PushNotificationNotSupported
| ErrorCode::UnsupportedOperation
| ErrorCode::VersionNotSupported => tonic::Code::Unimplemented,
ErrorCode::InvalidAgentResponse | ErrorCode::InternalError | _ => tonic::Code::Internal,
};
if let Some(reason) = a2a_err.code.a2a_reason() {
use tonic_types::StatusExt as _;
let mut details = tonic_types::ErrorDetails::new();
details.set_error_info(
reason,
a2a_protocol_types::error::A2A_ERROR_DOMAIN,
HashMap::<String, String>::new(),
);
return Status::with_error_details(code, a2a_err.message, details);
}
Status::new(code, a2a_err.message)
}
pub(super) async fn resolve_addr(
addr: impl tokio::net::ToSocketAddrs,
) -> std::io::Result<SocketAddr> {
tokio::net::lookup_host(addr).await?.next().ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::AddrNotAvailable,
"could not resolve address",
)
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::ServerError;
#[test]
fn task_not_found_maps_to_not_found() {
let status = server_error_to_status(&ServerError::TaskNotFound("t1".into()));
assert_eq!(status.code(), tonic::Code::NotFound);
}
#[test]
fn task_not_cancelable_maps_to_failed_precondition() {
let status = server_error_to_status(&ServerError::TaskNotCancelable("t1".into()));
assert_eq!(status.code(), tonic::Code::FailedPrecondition);
}
#[test]
fn invalid_params_maps_to_invalid_argument() {
let status = server_error_to_status(&ServerError::InvalidParams("bad".into()));
assert_eq!(status.code(), tonic::Code::InvalidArgument);
}
#[test]
fn method_not_found_maps_to_unimplemented() {
let status = server_error_to_status(&ServerError::MethodNotFound("Foo".into()));
assert_eq!(status.code(), tonic::Code::Unimplemented);
}
#[test]
fn internal_error_maps_to_internal() {
let status = server_error_to_status(&ServerError::Internal("oops".into()));
assert_eq!(status.code(), tonic::Code::Internal);
}
#[test]
fn a2a_error_status_carries_error_info_detail() {
use tonic_types::StatusExt as _;
let status = server_error_to_status(&ServerError::TaskNotFound("t1".into()));
let info = status
.get_details_error_info()
.expect("A2A error must carry google.rpc.ErrorInfo in status.details");
assert_eq!(info.reason, "TASK_NOT_FOUND");
assert_eq!(info.domain, a2a_protocol_types::error::A2A_ERROR_DOMAIN);
}
#[test]
fn every_a2a_reason_maps_to_error_info_detail() {
use tonic_types::StatusExt as _;
let cases: Vec<(ServerError, &str)> = vec![
(ServerError::TaskNotFound("t".into()), "TASK_NOT_FOUND"),
(
ServerError::TaskNotCancelable("t".into()),
"TASK_NOT_CANCELABLE",
),
(
ServerError::PushNotSupported,
"PUSH_NOTIFICATION_NOT_SUPPORTED",
),
(
ServerError::UnsupportedOperation("u".into()),
"UNSUPPORTED_OPERATION",
),
];
for (err, want_reason) in cases {
let status = server_error_to_status(&err);
let info = status
.get_details_error_info()
.unwrap_or_else(|| panic!("missing ErrorInfo for {want_reason}"));
assert_eq!(info.reason, want_reason);
assert_eq!(info.domain, "a2a-protocol.org");
}
}
#[test]
fn standard_errors_have_no_error_info_detail() {
use tonic_types::StatusExt as _;
for err in [
ServerError::Internal("oops".into()),
ServerError::InvalidParams("bad".into()),
ServerError::Overloaded("busy".into()),
] {
let status = server_error_to_status(&err);
assert!(
status.get_details_error_info().is_none(),
"standard error unexpectedly carried ErrorInfo: {err:?}"
);
}
}
#[test]
fn extract_metadata_ascii_keys() {
let mut meta = tonic::metadata::MetadataMap::new();
meta.insert("authorization", "Bearer token".parse().unwrap());
let map = extract_metadata(&meta);
assert_eq!(
map.get("authorization").map(String::as_str),
Some("Bearer token")
);
}
#[test]
fn extract_metadata_empty() {
let meta = tonic::metadata::MetadataMap::new();
let map = extract_metadata(&meta);
assert!(map.is_empty());
}
#[test]
fn validated_metadata_accepts_1x_and_absent() {
for version in [Some("1.0"), Some("1.5"), Some(""), None] {
let mut meta = tonic::metadata::MetadataMap::new();
if let Some(v) = version {
meta.insert("a2a-version", v.parse().unwrap());
}
assert!(
validated_metadata(&meta).is_ok(),
"version {version:?} must be accepted"
);
}
}
#[test]
fn validated_metadata_rejects_unsupported_versions() {
use tonic_types::StatusExt as _;
for version in ["0.3", "2.0", "not-a-version"] {
let mut meta = tonic::metadata::MetadataMap::new();
meta.insert("a2a-version", version.parse().unwrap());
let status =
validated_metadata(&meta).expect_err("unsupported version must be rejected");
assert_eq!(status.code(), tonic::Code::Unimplemented);
let info = status
.get_details_error_info()
.expect("version rejection must carry ErrorInfo");
assert_eq!(info.reason, "VERSION_NOT_SUPPORTED");
}
}
}