use std::collections::HashMap;
use a2a_protocol_types::message::Message;
use a2a_protocol_types::params::ListTasksParams;
use a2a_protocol_types::task::Task;
use crate::call_context::CallContext;
use crate::error::{ServerError, ServerResult};
use super::RequestHandler;
pub(super) fn validate_id(raw: &str, name: &str, max_length: usize) -> ServerResult<()> {
let trimmed = raw.trim();
if trimmed.is_empty() {
return Err(ServerError::InvalidParams(format!(
"{name} must not be empty or whitespace-only"
)));
}
if trimmed.len() > max_length {
return Err(ServerError::InvalidParams(format!(
"{name} exceeds maximum length (got {}, max {max_length})",
trimmed.len()
)));
}
Ok(())
}
pub(super) fn validate_metadata_object(
metadata: Option<&serde_json::Value>,
field: &str,
) -> ServerResult<()> {
if let Some(value) = metadata {
if !value.is_object() {
return Err(ServerError::InvalidParams(format!(
"{field} metadata must be a JSON object (got {}); non-object metadata \
is not representable across all A2A transports (gRPC google.protobuf.Struct)",
json_kind(value)
)));
}
}
Ok(())
}
const fn json_kind(value: &serde_json::Value) -> &'static str {
match value {
serde_json::Value::Null => "null",
serde_json::Value::Bool(_) => "boolean",
serde_json::Value::Number(_) => "number",
serde_json::Value::String(_) => "string",
serde_json::Value::Array(_) => "array",
serde_json::Value::Object(_) => "object",
}
}
pub(super) fn build_call_context(
method: &str,
headers: Option<&HashMap<String, String>>,
) -> CallContext {
let mut ctx = CallContext::new(method);
if let Some(h) = headers {
let extensions = parse_extensions_header(h);
if !extensions.is_empty() {
ctx = ctx.with_extensions(extensions);
}
ctx = ctx.with_http_headers(h.clone());
}
ctx
}
pub(super) fn parse_extensions_header(headers: &HashMap<String, String>) -> Vec<String> {
headers
.get("a2a-extensions")
.map(|v| {
v.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_owned)
.collect()
})
.unwrap_or_default()
}
pub(super) fn truncate_history(history: Option<Vec<Message>>, n: u32) -> Option<Vec<Message>> {
if n == 0 {
return None;
}
let mut msgs = history?;
let excess = msgs.len().saturating_sub(n as usize);
msgs.drain(..excess);
Some(msgs)
}
impl RequestHandler {
const CONTEXT_LOOKUP_PAGE_SIZE: u32 = 10;
pub(crate) async fn find_task_by_context(
&self,
context_id: &str,
) -> ServerResult<Option<Task>> {
if context_id.len() > self.limits.max_id_length {
return Ok(None);
}
let tenant = crate::store::tenant::TenantContext::current();
let tenant_param = if tenant.is_empty() {
None
} else {
Some(tenant)
};
let params = ListTasksParams {
tenant: tenant_param,
context_id: Some(context_id.to_owned()),
status: None,
page_size: Some(Self::CONTEXT_LOOKUP_PAGE_SIZE),
page_token: None,
status_timestamp_after: None,
include_artifacts: None,
history_length: None,
};
let resp = self.task_store.list(¶ms).await?;
let mut terminal_fallback: Option<Task> = None;
for task in resp.tasks {
if !task.status.state.is_terminal() {
return Ok(Some(task));
}
if terminal_fallback.is_none() {
terminal_fallback = Some(task);
}
}
Ok(terminal_fallback)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn history(len: usize) -> Vec<Message> {
use a2a_protocol_types::message::{MessageId, MessageRole, Part};
(0..len)
.map(|i| Message {
id: MessageId::new(format!("h{i}")),
role: MessageRole::User,
parts: vec![Part::text("x")],
context_id: None,
task_id: None,
reference_task_ids: None,
extensions: None,
metadata: None,
})
.collect()
}
fn kept(len: usize, n: u32) -> Option<Vec<String>> {
truncate_history(Some(history(len)), n)
.map(|msgs| msgs.into_iter().map(|m| m.id.0).collect())
}
#[test]
fn truncate_history_zero_omits_rather_than_emptying() {
assert_eq!(kept(3, 0), None, "n=0 must omit history entirely");
assert_eq!(kept(0, 0), None);
}
#[test]
fn truncate_history_keeps_the_most_recent() {
assert_eq!(kept(6, 2), Some(vec!["h4".into(), "h5".into()]));
assert_eq!(kept(4, 1), Some(vec!["h3".into()]));
}
#[test]
fn truncate_history_at_or_below_the_limit_keeps_everything() {
assert_eq!(
kept(3, 3),
Some(vec!["h0".into(), "h1".into(), "h2".into()])
);
assert_eq!(kept(2, 5), Some(vec!["h0".into(), "h1".into()]));
}
#[test]
fn truncate_history_distinguishes_absent_from_empty() {
assert_eq!(truncate_history(None, 3), None);
assert_eq!(kept(0, 3), Some(vec![]));
}
#[test]
fn validate_id_accepts_normal_id() {
assert!(
validate_id("task-123", "task_id", 1024).is_ok(),
"a normal short ID should be accepted"
);
}
#[test]
fn validate_id_rejects_empty_string() {
let err = validate_id("", "task_id", 1024).unwrap_err();
assert!(
matches!(err, ServerError::InvalidParams(ref msg) if msg.contains("empty")),
"empty string should be rejected with InvalidParams: {err:?}"
);
}
#[test]
fn validate_id_rejects_whitespace_only() {
let err = validate_id(" \t\n ", "context_id", 1024).unwrap_err();
assert!(
matches!(err, ServerError::InvalidParams(ref msg) if msg.contains("empty")),
"whitespace-only string should be rejected: {err:?}"
);
}
#[test]
fn validate_id_rejects_exceeding_max_length() {
let long_id = "a".repeat(2000);
let err = validate_id(&long_id, "task_id", 1024).unwrap_err();
assert!(
matches!(err, ServerError::InvalidParams(ref msg) if msg.contains("maximum length")),
"overly long ID should be rejected: {err:?}"
);
}
#[test]
fn validate_id_accepts_exactly_max_length() {
let exact = "b".repeat(128);
assert!(
validate_id(&exact, "task_id", 128).is_ok(),
"ID at exactly max length should be accepted"
);
}
#[test]
fn validate_id_trims_before_length_check() {
assert!(
validate_id(" abc ", "id", 3).is_ok(),
"trimmed length (3) should pass a max of 3"
);
}
#[test]
fn validate_id_includes_field_name_in_error() {
let err = validate_id("", "my_field", 1024).unwrap_err();
assert!(
matches!(err, ServerError::InvalidParams(ref msg) if msg.contains("my_field")),
"error message should contain the field name: {err:?}"
);
}
#[test]
fn build_call_context_without_headers() {
let ctx = build_call_context("message/send", None);
assert_eq!(ctx.method(), "message/send", "method should be set");
assert!(
ctx.http_headers().is_empty(),
"headers should be empty when None is passed"
);
}
#[test]
fn build_call_context_with_headers() {
let mut headers = HashMap::new();
headers.insert("authorization".to_owned(), "Bearer tok".to_owned());
headers.insert("x-request-id".to_owned(), "req-99".to_owned());
let ctx = build_call_context("tasks/get", Some(&headers));
assert_eq!(ctx.method(), "tasks/get");
assert_eq!(
ctx.http_headers().get("authorization").map(String::as_str),
Some("Bearer tok"),
"headers should be cloned into the context"
);
assert_eq!(
ctx.http_headers().get("x-request-id").map(String::as_str),
Some("req-99"),
);
}
#[test]
fn build_call_context_with_empty_headers_map() {
let headers = HashMap::new();
let ctx = build_call_context("test", Some(&headers));
assert!(
ctx.http_headers().is_empty(),
"an empty map should result in empty headers"
);
}
#[test]
fn build_call_context_parses_extensions_header() {
let mut headers = HashMap::new();
headers.insert(
"a2a-extensions".to_owned(),
"https://example.com/ext/geo/v1, https://standards.org/ext/cite/v1".to_owned(),
);
let ctx = build_call_context("message/send", Some(&headers));
assert_eq!(
ctx.extensions(),
&[
"https://example.com/ext/geo/v1".to_owned(),
"https://standards.org/ext/cite/v1".to_owned(),
],
"comma-separated extension URIs must be parsed and trimmed"
);
}
#[test]
fn build_call_context_no_extensions_header_is_empty() {
let mut headers = HashMap::new();
headers.insert("authorization".to_owned(), "Bearer tok".to_owned());
let ctx = build_call_context("message/send", Some(&headers));
assert!(ctx.extensions().is_empty());
}
#[test]
fn parse_extensions_header_drops_empty_segments() {
let mut headers = HashMap::new();
headers.insert(
"a2a-extensions".to_owned(),
" ,https://example.com/ext/v1,, ".to_owned(),
);
assert_eq!(
parse_extensions_header(&headers),
vec!["https://example.com/ext/v1".to_owned()],
"whitespace-only and empty segments must be dropped"
);
headers.insert("a2a-extensions".to_owned(), " ".to_owned());
assert!(
parse_extensions_header(&headers).is_empty(),
"a blank header value yields no extensions"
);
}
mod find_task_by_context_tests {
use a2a_protocol_types::task::{ContextId, Task, TaskId, TaskState, TaskStatus};
use crate::agent_executor;
use crate::builder::RequestHandlerBuilder;
use crate::handler::limits::HandlerLimits;
struct DummyExecutor;
agent_executor!(DummyExecutor, |_ctx, _queue| async { Ok(()) });
#[tokio::test]
async fn context_id_too_long_returns_none_without_querying_the_store() {
let handler = RequestHandlerBuilder::new(DummyExecutor)
.with_handler_limits(HandlerLimits::default().with_max_id_length(10))
.build()
.unwrap();
let long_id = "a".repeat(11);
handler
.task_store
.save(&make_task("t-long", &long_id, TaskState::Working))
.await
.unwrap();
let result = handler.find_task_by_context(&long_id).await.unwrap();
assert!(
result.is_none(),
"context_id longer than max_id_length must be rejected, but the \
store's matching task came back: {result:?}"
);
}
#[tokio::test]
async fn context_id_at_exactly_max_length_is_still_looked_up() {
let handler = RequestHandlerBuilder::new(DummyExecutor)
.with_handler_limits(HandlerLimits::default().with_max_id_length(10))
.build()
.unwrap();
let exact_id = "a".repeat(10);
handler
.task_store
.save(&make_task("t-exact", &exact_id, TaskState::Working))
.await
.unwrap();
let found = handler
.find_task_by_context(&exact_id)
.await
.unwrap()
.expect("a context_id of exactly max_id_length is within the limit");
assert_eq!(found.id.0, "t-exact");
}
#[tokio::test]
async fn context_id_within_limit_returns_none_for_missing() {
let handler = RequestHandlerBuilder::new(DummyExecutor)
.with_handler_limits(HandlerLimits::default().with_max_id_length(100))
.build()
.unwrap();
let result = handler
.find_task_by_context("no-such-context")
.await
.unwrap();
assert!(
result.is_none(),
"find_task_by_context should return None when no task matches the context"
);
}
fn make_task(id: &str, context_id: &str, state: TaskState) -> Task {
Task {
id: TaskId::new(id.to_owned()),
context_id: ContextId::new(context_id),
status: TaskStatus::new(state),
history: None,
artifacts: None,
metadata: None,
}
}
#[tokio::test]
async fn prefers_non_terminal_over_terminal_task() {
let handler = RequestHandlerBuilder::new(DummyExecutor)
.with_handler_limits(HandlerLimits::default().with_max_id_length(100))
.build()
.unwrap();
handler
.task_store
.save(&make_task("aaa-completed", "ctx-1", TaskState::Completed))
.await
.unwrap();
handler
.task_store
.save(&make_task("bbb-working", "ctx-1", TaskState::Working))
.await
.unwrap();
let result = handler.find_task_by_context("ctx-1").await.unwrap();
assert!(result.is_some(), "should find a task");
let task = result.unwrap();
assert_eq!(
task.id.0, "bbb-working",
"should prefer the non-terminal (Working) task over the terminal (Completed) one"
);
}
#[tokio::test]
async fn returns_terminal_task_when_no_non_terminal_exists() {
let handler = RequestHandlerBuilder::new(DummyExecutor)
.with_handler_limits(HandlerLimits::default().with_max_id_length(100))
.build()
.unwrap();
handler
.task_store
.save(&make_task("task-done", "ctx-2", TaskState::Completed))
.await
.unwrap();
let result = handler.find_task_by_context("ctx-2").await.unwrap();
assert!(result.is_some(), "should still return a terminal task");
assert_eq!(result.unwrap().id.0, "task-done");
}
#[tokio::test]
async fn returns_first_non_terminal_when_multiple_exist() {
let handler = RequestHandlerBuilder::new(DummyExecutor)
.with_handler_limits(HandlerLimits::default().with_max_id_length(100))
.build()
.unwrap();
handler
.task_store
.save(&make_task("aaa-failed", "ctx-3", TaskState::Failed))
.await
.unwrap();
handler
.task_store
.save(&make_task("bbb-submitted", "ctx-3", TaskState::Submitted))
.await
.unwrap();
handler
.task_store
.save(&make_task("ccc-working", "ctx-3", TaskState::Working))
.await
.unwrap();
let result = handler.find_task_by_context("ctx-3").await.unwrap();
let task = result.unwrap();
assert!(
!task.status.state.is_terminal(),
"should return a non-terminal task, got {:?}",
task.status.state
);
assert_eq!(
task.id.0, "ccc-working",
"should return the most-recently-updated non-terminal task"
);
}
}
}