use std::collections::HashMap;
use std::fmt;
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
use serde::de::{self, IgnoredAny, MapAccess, Visitor};
use serde::ser::SerializeMap;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use serde_json::Value;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Role {
#[default]
Unspecified,
User,
Agent,
}
impl Role {
const fn as_wire(self) -> &'static str {
match self {
Self::Unspecified => "ROLE_UNSPECIFIED",
Self::User => "ROLE_USER",
Self::Agent => "ROLE_AGENT",
}
}
}
impl Serialize for Role {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(self.as_wire())
}
}
impl<'de> Deserialize<'de> for Role {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(match String::deserialize(deserializer)?.as_str() {
"ROLE_USER" => Self::User,
"ROLE_AGENT" => Self::Agent,
_ => Self::Unspecified,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum TaskState {
#[default]
Unspecified,
Submitted,
Working,
Completed,
Failed,
Canceled,
InputRequired,
Rejected,
AuthRequired,
}
impl TaskState {
const fn as_wire(self) -> &'static str {
match self {
Self::Unspecified => "TASK_STATE_UNSPECIFIED",
Self::Submitted => "TASK_STATE_SUBMITTED",
Self::Working => "TASK_STATE_WORKING",
Self::Completed => "TASK_STATE_COMPLETED",
Self::Failed => "TASK_STATE_FAILED",
Self::Canceled => "TASK_STATE_CANCELED",
Self::InputRequired => "TASK_STATE_INPUT_REQUIRED",
Self::Rejected => "TASK_STATE_REJECTED",
Self::AuthRequired => "TASK_STATE_AUTH_REQUIRED",
}
}
#[must_use]
pub const fn is_terminal(self) -> bool {
matches!(
self,
Self::Completed | Self::Failed | Self::Canceled | Self::Rejected
)
}
}
impl Serialize for TaskState {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(self.as_wire())
}
}
impl<'de> Deserialize<'de> for TaskState {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(match String::deserialize(deserializer)?.as_str() {
"TASK_STATE_SUBMITTED" => Self::Submitted,
"TASK_STATE_WORKING" => Self::Working,
"TASK_STATE_COMPLETED" => Self::Completed,
"TASK_STATE_FAILED" => Self::Failed,
"TASK_STATE_CANCELED" => Self::Canceled,
"TASK_STATE_INPUT_REQUIRED" => Self::InputRequired,
"TASK_STATE_REJECTED" => Self::Rejected,
"TASK_STATE_AUTH_REQUIRED" => Self::AuthRequired,
_ => Self::Unspecified,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PartContent {
Text(String),
Raw(Vec<u8>),
Url(String),
Data(Value),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Part {
pub content: PartContent,
pub filename: Option<String>,
pub media_type: Option<String>,
}
impl Part {
#[must_use]
pub fn text(text: impl Into<String>) -> Self {
Self {
content: PartContent::Text(text.into()),
filename: None,
media_type: None,
}
}
#[must_use]
pub fn as_text(&self) -> Option<&str> {
match &self.content {
PartContent::Text(t) => Some(t),
_ => None,
}
}
}
impl Serialize for Part {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let mut map = serializer.serialize_map(None)?;
match &self.content {
PartContent::Text(t) => map.serialize_entry("text", t)?,
PartContent::Raw(r) => map.serialize_entry("raw", &BASE64.encode(r))?,
PartContent::Url(u) => map.serialize_entry("url", u)?,
PartContent::Data(d) => map.serialize_entry("data", d)?,
}
if let Some(filename) = &self.filename {
map.serialize_entry("filename", filename)?;
}
if let Some(media_type) = &self.media_type {
map.serialize_entry("mediaType", media_type)?;
}
map.end()
}
}
impl<'de> Deserialize<'de> for Part {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct PartVisitor;
impl<'de> Visitor<'de> for PartVisitor {
type Value = Part;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("an A2A Part object with one of text/raw/url/data")
}
fn visit_map<M: MapAccess<'de>>(self, mut map: M) -> Result<Part, M::Error> {
let mut content: Option<PartContent> = None;
let mut filename = None;
let mut media_type = None;
while let Some(key) = map.next_key::<String>()? {
match key.as_str() {
"text" => content = Some(PartContent::Text(map.next_value()?)),
"raw" => {
let b64: String = map.next_value()?;
let bytes = BASE64.decode(b64).map_err(de::Error::custom)?;
content = Some(PartContent::Raw(bytes));
}
"url" => content = Some(PartContent::Url(map.next_value()?)),
"data" => content = Some(PartContent::Data(map.next_value()?)),
"filename" => filename = Some(map.next_value()?),
"mediaType" => media_type = Some(map.next_value()?),
_ => {
map.next_value::<IgnoredAny>()?;
}
}
}
let content = content.ok_or_else(|| {
de::Error::custom("Part has no content key (text/raw/url/data)")
})?;
Ok(Part {
content,
filename,
media_type,
})
}
}
deserializer.deserialize_map(PartVisitor)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
#[allow(clippy::struct_field_names)]
pub struct Message {
#[serde(default)]
pub message_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub context_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub task_id: Option<String>,
#[serde(default)]
pub role: Role,
#[serde(default)]
pub parts: Vec<Part>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, Value>>,
}
impl Message {
#[must_use]
pub fn text(&self) -> String {
self.parts
.iter()
.filter_map(Part::as_text)
.collect::<Vec<_>>()
.join("\n")
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct TaskStatus {
#[serde(default)]
pub state: TaskState,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub message: Option<Message>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timestamp: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Task {
#[serde(default)]
pub id: String,
#[serde(default)]
pub context_id: String,
#[serde(default)]
pub status: TaskStatus,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub artifacts: Option<Vec<Artifact>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub history: Option<Vec<Message>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, Value>>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
#[allow(clippy::struct_field_names)]
pub struct Artifact {
#[serde(default)]
pub artifact_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default)]
pub parts: Vec<Part>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, Value>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SendMessageResponse {
Task(Task),
Message(Message),
}
impl Serialize for SendMessageResponse {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let mut map = serializer.serialize_map(Some(1))?;
match self {
Self::Task(task) => map.serialize_entry("task", task)?,
Self::Message(message) => map.serialize_entry("message", message)?,
}
map.end()
}
}
impl<'de> Deserialize<'de> for SendMessageResponse {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let value = Value::deserialize(deserializer)?;
if let Some(task) = value.get("task") {
return Task::deserialize(task.clone())
.map(Self::Task)
.map_err(de::Error::custom);
}
if let Some(message) = value.get("message") {
return Message::deserialize(message.clone())
.map(Self::Message)
.map_err(de::Error::custom);
}
Err(de::Error::custom(
"SendMessageResponse has neither `task` nor `message`",
))
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum StreamResponse {
Task(Task),
StatusUpdate(TaskStatusUpdateEvent),
ArtifactUpdate(TaskArtifactUpdateEvent),
}
impl Serialize for StreamResponse {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let mut map = serializer.serialize_map(Some(1))?;
match self {
Self::Task(task) => map.serialize_entry("task", task)?,
Self::StatusUpdate(update) => map.serialize_entry("statusUpdate", update)?,
Self::ArtifactUpdate(update) => map.serialize_entry("artifactUpdate", update)?,
}
map.end()
}
}
impl<'de> Deserialize<'de> for StreamResponse {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let value = Value::deserialize(deserializer)?;
if let Some(task) = value.get("task") {
return Task::deserialize(task.clone())
.map(Self::Task)
.map_err(de::Error::custom);
}
if let Some(update) = value.get("statusUpdate") {
return TaskStatusUpdateEvent::deserialize(update.clone())
.map(Self::StatusUpdate)
.map_err(de::Error::custom);
}
if let Some(update) = value.get("artifactUpdate") {
return TaskArtifactUpdateEvent::deserialize(update.clone())
.map(Self::ArtifactUpdate)
.map_err(de::Error::custom);
}
Err(de::Error::custom(
"StreamResponse has none of `task`/`statusUpdate`/`artifactUpdate`",
))
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
#[allow(clippy::struct_field_names)]
pub struct TaskStatusUpdateEvent {
pub task_id: String,
pub context_id: String,
pub status: TaskStatus,
#[serde(rename = "final")]
pub is_final: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
#[allow(clippy::struct_field_names)]
pub struct TaskArtifactUpdateEvent {
pub task_id: String,
pub context_id: String,
pub artifact: Artifact,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub append: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub last_chunk: Option<bool>,
}
#[cfg(test)]
mod tests {
#![allow(clippy::pedantic, clippy::nursery, missing_docs)]
use super::*;
use serde_json::json;
#[test]
fn role_wire_is_screaming_snake() {
assert_eq!(
serde_json::to_string(&Role::User).unwrap(),
r#""ROLE_USER""#
);
assert_eq!(
serde_json::to_string(&Role::Agent).unwrap(),
r#""ROLE_AGENT""#
);
assert_eq!(
serde_json::from_str::<Role>(r#""ROLE_AGENT""#).unwrap(),
Role::Agent
);
assert_eq!(
serde_json::from_str::<Role>(r#""ROLE_FUTURE""#).unwrap(),
Role::Unspecified
);
}
#[test]
fn task_state_wire_is_screaming_snake() {
assert_eq!(
serde_json::to_string(&TaskState::Working).unwrap(),
r#""TASK_STATE_WORKING""#
);
assert_eq!(
serde_json::to_string(&TaskState::Completed).unwrap(),
r#""TASK_STATE_COMPLETED""#
);
assert_eq!(
serde_json::to_string(&TaskState::Canceled).unwrap(),
r#""TASK_STATE_CANCELED""#
);
assert_eq!(
serde_json::to_string(&TaskState::InputRequired).unwrap(),
r#""TASK_STATE_INPUT_REQUIRED""#
);
assert_eq!(
serde_json::from_str::<TaskState>(r#""TASK_STATE_INPUT_REQUIRED""#).unwrap(),
TaskState::InputRequired
);
}
#[test]
fn terminal_states() {
assert!(TaskState::Completed.is_terminal());
assert!(TaskState::Canceled.is_terminal());
assert!(!TaskState::Working.is_terminal());
assert!(!TaskState::InputRequired.is_terminal());
}
#[test]
fn part_is_field_presence_no_kind() {
let text = Part::text("hello");
assert_eq!(
serde_json::to_value(&text).unwrap(),
json!({ "text": "hello" })
);
let back: Part = serde_json::from_value(json!({ "text": "hello" })).unwrap();
assert_eq!(back, text);
let raw = Part {
content: PartContent::Raw(vec![0, 1, 2, 3]),
filename: Some("x.bin".to_owned()),
media_type: Some("application/octet-stream".to_owned()),
};
assert_eq!(
serde_json::to_value(&raw).unwrap(),
json!({ "raw": "AAECAw==", "filename": "x.bin", "mediaType": "application/octet-stream" })
);
let data = Part {
content: PartContent::Data(json!({ "k": 1 })),
filename: None,
media_type: None,
};
assert_eq!(
serde_json::to_value(&data).unwrap(),
json!({ "data": { "k": 1 } })
);
assert!(serde_json::from_value::<Part>(json!({ "filename": "x" })).is_err());
}
#[test]
fn message_is_camel_case_no_kind() {
let msg = Message {
message_id: "m1".to_owned(),
context_id: Some("c1".to_owned()),
task_id: None,
role: Role::User,
parts: vec![Part::text("hi")],
metadata: None,
};
let value = serde_json::to_value(&msg).unwrap();
assert_eq!(
value,
json!({
"messageId": "m1",
"contextId": "c1",
"role": "ROLE_USER",
"parts": [{ "text": "hi" }]
})
);
assert!(value.get("kind").is_none(), "v1.0 messages carry no `kind`");
assert_eq!(serde_json::from_value::<Message>(value).unwrap(), msg);
}
#[test]
fn task_and_status_shape() {
let task = Task {
id: "t1".to_owned(),
context_id: "c1".to_owned(),
status: TaskStatus {
state: TaskState::Completed,
message: Some(Message {
message_id: "r1".to_owned(),
context_id: Some("c1".to_owned()),
task_id: Some("t1".to_owned()),
role: Role::Agent,
parts: vec![Part::text("the answer is 42")],
metadata: None,
}),
timestamp: None,
},
artifacts: None,
history: None,
metadata: None,
};
let value = serde_json::to_value(&task).unwrap();
assert_eq!(value["id"], "t1");
assert_eq!(value["contextId"], "c1");
assert_eq!(value["status"]["state"], "TASK_STATE_COMPLETED");
assert_eq!(
value["status"]["message"]["parts"][0]["text"],
"the answer is 42"
);
assert!(value.get("kind").is_none(), "v1.0 tasks carry no `kind`");
assert_eq!(serde_json::from_value::<Task>(value).unwrap(), task);
}
#[test]
fn send_message_response_is_wrapped_oneof() {
let task = Task {
id: "t1".to_owned(),
context_id: "c1".to_owned(),
status: TaskStatus {
state: TaskState::Completed,
message: None,
timestamp: None,
},
artifacts: None,
history: None,
metadata: None,
};
let resp = SendMessageResponse::Task(task.clone());
let value = serde_json::to_value(&resp).unwrap();
assert_eq!(value["task"]["id"], "t1");
assert!(value.get("message").is_none());
assert_eq!(
serde_json::from_value::<SendMessageResponse>(value).unwrap(),
resp
);
}
#[test]
fn status_update_event_is_camel_case_with_final_keyword() {
let update = TaskStatusUpdateEvent {
task_id: "t1".to_owned(),
context_id: "c1".to_owned(),
status: TaskStatus {
state: TaskState::Working,
message: None,
timestamp: None,
},
is_final: false,
};
let value = serde_json::to_value(&update).unwrap();
assert_eq!(
value,
json!({
"taskId": "t1",
"contextId": "c1",
"status": { "state": "TASK_STATE_WORKING" },
"final": false,
})
);
assert_eq!(
serde_json::from_value::<TaskStatusUpdateEvent>(value).unwrap(),
update
);
}
#[test]
fn artifact_update_event_omits_absent_append_and_last_chunk() {
let update = TaskArtifactUpdateEvent {
task_id: "t1".to_owned(),
context_id: "c1".to_owned(),
artifact: Artifact {
artifact_id: "a1".to_owned(),
name: None,
description: None,
parts: vec![Part::text("chunk one")],
metadata: None,
},
append: None,
last_chunk: None,
};
let value = serde_json::to_value(&update).unwrap();
assert_eq!(
value,
json!({
"taskId": "t1",
"contextId": "c1",
"artifact": { "artifactId": "a1", "parts": [{ "text": "chunk one" }] },
})
);
assert!(value.get("append").is_none());
assert!(value.get("lastChunk").is_none());
assert_eq!(
serde_json::from_value::<TaskArtifactUpdateEvent>(value).unwrap(),
update
);
let next = TaskArtifactUpdateEvent {
append: Some(true),
..update
};
assert_eq!(serde_json::to_value(&next).unwrap()["append"], true);
}
#[test]
fn stream_response_is_wrapped_oneof_over_all_three_variants() {
let task = Task {
id: "t1".to_owned(),
context_id: "c1".to_owned(),
status: TaskStatus::default(),
artifacts: None,
history: None,
metadata: None,
};
let task_value = serde_json::to_value(StreamResponse::Task(task.clone())).unwrap();
assert_eq!(task_value["task"]["id"], "t1");
assert!(task_value.get("statusUpdate").is_none());
assert!(task_value.get("artifactUpdate").is_none());
assert_eq!(
serde_json::from_value::<StreamResponse>(task_value).unwrap(),
StreamResponse::Task(task)
);
let status = TaskStatusUpdateEvent {
task_id: "t1".to_owned(),
context_id: "c1".to_owned(),
status: TaskStatus {
state: TaskState::Completed,
message: None,
timestamp: None,
},
is_final: true,
};
let status_value =
serde_json::to_value(StreamResponse::StatusUpdate(status.clone())).unwrap();
assert_eq!(status_value["statusUpdate"]["final"], true);
assert_eq!(
serde_json::from_value::<StreamResponse>(status_value).unwrap(),
StreamResponse::StatusUpdate(status)
);
let artifact_update = TaskArtifactUpdateEvent {
task_id: "t1".to_owned(),
context_id: "c1".to_owned(),
artifact: Artifact {
artifact_id: "a1".to_owned(),
name: None,
description: None,
parts: vec![Part::text("hi")],
metadata: None,
},
append: None,
last_chunk: None,
};
let artifact_value =
serde_json::to_value(StreamResponse::ArtifactUpdate(artifact_update.clone())).unwrap();
assert_eq!(artifact_value["artifactUpdate"]["taskId"], "t1");
assert_eq!(
serde_json::from_value::<StreamResponse>(artifact_value).unwrap(),
StreamResponse::ArtifactUpdate(artifact_update)
);
assert!(
serde_json::from_value::<StreamResponse>(json!({})).is_err(),
"an empty object matches none of the three variants"
);
}
#[test]
fn decodes_protojson_omitted_default_fields() {
let task: Task = serde_json::from_value(json!({ "id": "t" })).unwrap();
assert_eq!(task.context_id, "");
assert_eq!(task.status.state, TaskState::Unspecified);
assert!(task.artifacts.is_none());
let msg: Message = serde_json::from_value(json!({ "messageId": "m" })).unwrap();
assert_eq!(msg.role, Role::Unspecified);
assert!(msg.parts.is_empty());
}
}