use crate::auth::{Identity, RequestContext};
use crate::envelope::{A2aHeaders, A2aMethod, JsonRpcRequest, JsonRpcResponse};
use crate::error::A2aError;
use crate::handler::A2aHandler;
use crate::task_store::A2aTaskStore;
use crate::types::{
CancelTaskParams, Message, Part, Role, SendMessageParams, SendMessageResult,
SubscribeToTaskParams, Task, TaskStatus,
};
use klieo_bus_memory::MemoryKv;
use klieo_core::Headers;
use serde_json::json;
use std::borrow::Cow;
use std::fmt;
use std::time::Duration;
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum CaseStatus {
Pass,
Fail,
Skip(Cow<'static, str>),
}
impl CaseStatus {
fn label(&self) -> &'static str {
match self {
CaseStatus::Pass => "PASS",
CaseStatus::Fail => "FAIL",
CaseStatus::Skip(_) => "SKIP",
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct ConformanceCase {
pub id: &'static str,
pub clause: &'static str,
pub status: CaseStatus,
pub detail: Option<String>,
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct ConformanceReport {
pub passed: u32,
pub failed: u32,
pub cases: Vec<ConformanceCase>,
}
impl ConformanceReport {
pub fn skipped(&self) -> u32 {
self.cases
.iter()
.filter(|c| matches!(c.status, CaseStatus::Skip(_)))
.count() as u32
}
fn record(&mut self, case: ConformanceCase) {
match case.status {
CaseStatus::Pass => self.passed += 1,
CaseStatus::Fail => self.failed += 1,
CaseStatus::Skip(_) => {}
}
self.cases.push(case);
}
}
impl fmt::Display for ConformanceReport {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let total = self.passed + self.failed;
writeln!(
f,
"{}/{} passed (skips: {})",
self.passed,
total,
self.skipped()
)?;
writeln!(f, "{:<6} {:<6} {:<14} detail", "case", "status", "clause")?;
for c in &self.cases {
let detail = c.detail.as_deref().unwrap_or("—");
writeln!(
f,
"{:<6} {:<6} {:<14} {}",
c.id,
c.status.label(),
c.clause,
detail
)?;
}
Ok(())
}
}
fn pass(id: &'static str, clause: &'static str) -> ConformanceCase {
ConformanceCase {
id,
clause,
status: CaseStatus::Pass,
detail: None,
}
}
fn fail(id: &'static str, clause: &'static str, detail: impl Into<String>) -> ConformanceCase {
ConformanceCase {
id,
clause,
status: CaseStatus::Fail,
detail: Some(detail.into()),
}
}
fn skip(
id: &'static str,
clause: &'static str,
reason: impl Into<Cow<'static, str>>,
) -> ConformanceCase {
let reason = reason.into();
let detail = reason.to_string();
ConformanceCase {
id,
clause,
status: CaseStatus::Skip(reason),
detail: Some(detail),
}
}
pub async fn run_conformance_suite<H: A2aHandler + Send + Sync>(handler: &H) -> ConformanceReport {
let mut report = ConformanceReport {
cases: Vec::with_capacity(12),
..ConformanceReport::default()
};
report.record(case_c01_default_a2a_version().await);
report.record(case_c02_method_names_parse().await);
report.record(case_c03_unknown_method_rejects().await);
report.record(case_c04_id_null_accepted().await);
report.record(case_c05_id_int_and_string_round_trip().await);
report.record(case_c06_error_code_mapping().await);
report.record(case_c07_idempotency_key_header_round_trip().await);
report.record(case_c08_idempotency_replay().await);
report.record(case_c09_streaming_subscription_in_order().await);
report.record(case_c10_cancel_completes_with_canceled_status(handler).await);
report.record(case_c11_task_store_put_get_round_trip().await);
report.record(case_c12_malformed_json_is_parse_error().await);
report
}
async fn case_c01_default_a2a_version() -> ConformanceCase {
let h = Headers::new();
let decoded = A2aHeaders::decode_from(&h);
if decoded.a2a_version == "1.0" {
pass("C-01", "§8")
} else {
fail(
"C-01",
"§8",
format!(
"A2A-Version default was `{}`, expected `1.0`",
decoded.a2a_version
),
)
}
}
async fn case_c02_method_names_parse() -> ConformanceCase {
let cases: [(&str, A2aMethod); 11] = [
("SendMessage", A2aMethod::SendMessage),
("SendStreamingMessage", A2aMethod::SendStreamingMessage),
("GetTask", A2aMethod::GetTask),
("ListTasks", A2aMethod::ListTasks),
("CancelTask", A2aMethod::CancelTask),
("SubscribeToTask", A2aMethod::SubscribeToTask),
(
"CreateTaskPushNotificationConfig",
A2aMethod::CreateTaskPushNotificationConfig,
),
(
"GetTaskPushNotificationConfig",
A2aMethod::GetTaskPushNotificationConfig,
),
(
"ListTaskPushNotificationConfigs",
A2aMethod::ListTaskPushNotificationConfigs,
),
(
"DeleteTaskPushNotificationConfig",
A2aMethod::DeleteTaskPushNotificationConfig,
),
("GetExtendedAgentCard", A2aMethod::GetExtendedAgentCard),
];
for (wire, expected) in cases {
match A2aMethod::from_str(wire) {
Ok(got) if got == expected => continue,
Ok(other) => {
return fail(
"C-02",
"§9.4",
format!("`{wire}` parsed as `{other:?}`, expected `{expected:?}`"),
);
}
Err(e) => {
return fail("C-02", "§9.4", format!("`{wire}` failed to parse: {e}"));
}
}
}
pass("C-02", "§9.4")
}
async fn case_c03_unknown_method_rejects() -> ConformanceCase {
match A2aMethod::from_str("Frobnicate") {
Err(A2aError::MethodNotFound(name)) if name == "Frobnicate" => pass("C-03", "§9.4"),
Err(other) => fail(
"C-03",
"§9.4",
format!("expected MethodNotFound, got {other:?}"),
),
Ok(m) => fail("C-03", "§9.4", format!("unknown name parsed as {m:?}")),
}
}
async fn case_c04_id_null_accepted() -> ConformanceCase {
let req = JsonRpcRequest {
jsonrpc: "2.0".into(),
id: serde_json::Value::Null,
method: "SendMessage".into(),
params: json!({}),
};
let v = match serde_json::to_value(&req) {
Ok(v) => v,
Err(e) => return fail("C-04", "§10", format!("encode: {e}")),
};
let back: Result<JsonRpcRequest, _> = serde_json::from_value(v);
match back {
Ok(r) if r.id.is_null() => pass("C-04", "§10"),
Ok(r) => fail("C-04", "§10", format!("id was {:?}, expected null", r.id)),
Err(e) => fail("C-04", "§10", format!("decode: {e}")),
}
}
async fn case_c05_id_int_and_string_round_trip() -> ConformanceCase {
for id in [json!(42_i64), json!("req-abc")] {
let req = JsonRpcRequest {
jsonrpc: "2.0".into(),
id: id.clone(),
method: "SendMessage".into(),
params: json!({}),
};
let v = match serde_json::to_value(&req) {
Ok(v) => v,
Err(e) => return fail("C-05", "§10", format!("encode {id:?}: {e}")),
};
let back: JsonRpcRequest = match serde_json::from_value(v) {
Ok(b) => b,
Err(e) => return fail("C-05", "§10", format!("decode {id:?}: {e}")),
};
if back.id != id {
return fail(
"C-05",
"§10",
format!("id round-trip mismatch: {:?} != {:?}", back.id, id),
);
}
}
pass("C-05", "§10")
}
async fn case_c06_error_code_mapping() -> ConformanceCase {
let id = json!(1);
let cases: [(A2aError, i32, &str); 4] = [
(
A2aError::MethodNotFound("X".into()),
-32601,
"MethodNotFound",
),
(A2aError::InvalidParams("X".into()), -32602, "InvalidParams"),
(A2aError::Server("X".into()), -32000, "Server"),
(A2aError::Unauthorized("X".into()), -32001, "Unauthorized"),
];
for (err, want, name) in cases {
let resp = err.to_json_rpc_error(id.clone());
let Some(payload) = resp.error else {
return fail("C-06", "§11", format!("{name}: missing error payload"));
};
if payload.code != want {
return fail(
"C-06",
"§11",
format!("{name}: code {} != expected {want}", payload.code),
);
}
}
let parse_err = match serde_json::from_str::<u32>("not json") {
Err(e) => A2aError::from(e),
Ok(_) => return fail("C-06", "§11", "cannot synthesise serde_json parse error"),
};
let resp = parse_err.to_json_rpc_error(id);
let Some(payload) = resp.error else {
return fail("C-06", "§11", "ParseError: missing error payload");
};
if payload.code != -32700 {
return fail(
"C-06",
"§11",
format!("ParseError: code {} != -32700", payload.code),
);
}
let server_resp = A2aError::Server("x".into()).to_json_rpc_error(json!(2));
let code = server_resp.error.unwrap().code;
if !(-32099..=-32000).contains(&code) {
return fail(
"C-06",
"§11",
format!("Server: code {code} outside -32000..-32099 band"),
);
}
pass("C-06", "§11")
}
async fn case_c07_idempotency_key_header_round_trip() -> ConformanceCase {
let mut h = Headers::new();
h.insert("Idempotency-Key".into(), "key-123".into());
h.insert("A2A-Version".into(), "1.0".into());
let decoded = A2aHeaders::decode_from(&h);
if decoded.a2a_version != "1.0" {
return fail(
"C-07",
"§12",
format!("a2a_version `{}` != `1.0`", decoded.a2a_version),
);
}
match h.get("Idempotency-Key") {
Some(v) if v == "key-123" => pass("C-07", "§12"),
Some(v) => fail(
"C-07",
"§12",
format!("Idempotency-Key round-trip mismatch: `{v}` != `key-123`"),
),
None => fail("C-07", "§12", "Idempotency-Key lost in Headers map"),
}
}
async fn case_c08_idempotency_replay() -> ConformanceCase {
skip("C-08", "§12", "idempotency ledger not yet implemented")
}
async fn case_c09_streaming_subscription_in_order() -> ConformanceCase {
let p = SubscribeToTaskParams { id: "t-1".into() };
let v = match serde_json::to_value(&p) {
Ok(v) => v,
Err(e) => return fail("C-09", "§13", format!("encode SubscribeToTaskParams: {e}")),
};
let back: SubscribeToTaskParams = match serde_json::from_value(v) {
Ok(b) => b,
Err(e) => return fail("C-09", "§13", format!("decode SubscribeToTaskParams: {e}")),
};
if back.id != "t-1" {
return fail("C-09", "§13", "SubscribeToTaskParams.id lost in round-trip");
}
skip(
"C-09",
"§13",
"chunked streaming delivery (SendStreamingMessage) not yet implemented",
)
}
async fn case_c10_cancel_completes_with_canceled_status<H: A2aHandler + Send + Sync>(
handler: &H,
) -> ConformanceCase {
let msg = Message {
messageId: "m-cancel".into(),
contextId: Some("ctx-cancel".into()),
taskId: None,
role: Role::User,
parts: vec![Part::Text {
content: "task body".into(),
metadata: None,
mediaType: None,
filename: None,
}],
metadata: None,
extensions: vec![],
referenceTaskIds: vec![],
};
let ctx = RequestContext::new(
A2aHeaders::decode_from(&klieo_core::Headers::new()),
Some(Identity::anonymous()),
);
let sent = match tokio::time::timeout(
Duration::from_secs(2),
handler.send_message(
&ctx,
SendMessageParams {
message: msg,
configuration: None,
},
),
)
.await
{
Ok(Ok(r)) => r,
Ok(Err(A2aError::MethodNotFound(_))) => {
return skip("C-10", "§13", "handler does not implement send_message")
}
Ok(Err(e)) => return fail("C-10", "§13", format!("send_message: {e}")),
Err(_) => return fail("C-10", "§13", "send_message exceeded 2s budget"),
};
let task_id = match sent {
SendMessageResult::Task(t) => t.id,
SendMessageResult::Message(_) => {
return skip(
"C-10",
"§13",
"handler returned a Message; no task to cancel",
)
}
};
match tokio::time::timeout(
Duration::from_secs(2),
handler.cancel_task(&ctx, CancelTaskParams { id: task_id }),
)
.await
{
Ok(Ok(task)) if matches!(task.status, TaskStatus::Canceled) => pass("C-10", "§13"),
Ok(Ok(task)) => fail(
"C-10",
"§13",
format!(
"cancel returned status {:?}, expected Canceled",
task.status
),
),
Ok(Err(A2aError::MethodNotFound(_))) => {
skip("C-10", "§13", "handler does not implement cancel_task")
}
Ok(Err(e)) => fail("C-10", "§13", format!("cancel_task: {e}")),
Err(_) => fail("C-10", "§13", "cancel_task exceeded 2s budget"),
}
}
async fn case_c11_task_store_put_get_round_trip() -> ConformanceCase {
let kv = MemoryKv::new();
let store = A2aTaskStore::new(std::sync::Arc::new(kv), "a2a.tasks".into());
let task = Task {
id: "conf-task-1".into(),
contextId: "conf-ctx-1".into(),
status: TaskStatus::Submitted,
artifacts: vec![],
history: vec![],
metadata: None,
};
if let Err(e) = store.put(&task).await {
return fail("C-11", "§14", format!("put: {e}"));
}
match store.get("conf-task-1").await {
Ok(Some(back)) if back.id == "conf-task-1" && back.contextId == "conf-ctx-1" => {
pass("C-11", "§14")
}
Ok(Some(back)) => fail(
"C-11",
"§14",
format!("got id={} ctx={}", back.id, back.contextId),
),
Ok(None) => fail("C-11", "§14", "task missing after put"),
Err(e) => fail("C-11", "§14", format!("get: {e}")),
}
}
async fn case_c12_malformed_json_is_parse_error() -> ConformanceCase {
let parse_result: Result<JsonRpcRequest, _> = serde_json::from_slice(b"not json {");
let err = match parse_result {
Err(e) => A2aError::from(e),
Ok(_) => {
return fail(
"C-12",
"§15",
"malformed body decoded as JsonRpcRequest unexpectedly",
)
}
};
let envelope: JsonRpcResponse = err.to_json_rpc_error(serde_json::Value::Null);
match envelope.error {
Some(p) if p.code == -32700 => pass("C-12", "§15"),
Some(p) => fail("C-12", "§15", format!("code {} != -32700", p.code)),
None => fail("C-12", "§15", "missing error payload"),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn case_status_labels() {
assert_eq!(CaseStatus::Pass.label(), "PASS");
assert_eq!(CaseStatus::Fail.label(), "FAIL");
assert_eq!(CaseStatus::Skip(Cow::Borrowed("x")).label(), "SKIP");
}
}