use chrono::{DateTime, Utc};
use hashbrown::HashMap;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "kebab-case")]
pub enum TaskState {
#[default]
Submitted,
Working,
InputRequired,
Completed,
Failed,
Canceled,
Rejected,
AuthRequired,
#[serde(other)]
Unknown,
}
impl TaskState {
pub fn is_terminal(&self) -> bool {
matches!(self, TaskState::Completed | TaskState::Failed | TaskState::Canceled | TaskState::Rejected)
}
pub fn is_cancelable(&self) -> bool {
matches!(self, TaskState::Submitted | TaskState::Working | TaskState::InputRequired)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TaskStatus {
pub state: TaskState,
#[serde(skip_serializing_if = "Option::is_none")]
pub message: Option<Message>,
pub timestamp: DateTime<Utc>,
}
impl TaskStatus {
pub fn new(state: TaskState) -> Self {
Self { state, message: None, timestamp: Utc::now() }
}
pub fn with_message(state: TaskState, message: Message) -> Self {
Self {
state,
message: Some(message),
timestamp: Utc::now(),
}
}
}
impl Default for TaskStatus {
fn default() -> Self {
Self::new(TaskState::Submitted)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum MessageRole {
User,
Agent,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Message {
pub role: MessageRole,
pub parts: Vec<Part>,
#[serde(skip_serializing_if = "Option::is_none")]
pub message_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub task_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub context_id: Option<String>,
#[serde(skip_serializing_if = "Vec::is_empty", default)]
pub reference_task_ids: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, serde_json::Value>>,
}
impl Message {
pub fn new(role: MessageRole, parts: Vec<Part>) -> Self {
Self {
role,
parts,
message_id: None,
task_id: None,
context_id: None,
reference_task_ids: Vec::new(),
metadata: None,
}
}
pub fn agent_text(text: impl Into<String>) -> Self {
Self::new(MessageRole::Agent, vec![Part::text(text)])
}
pub fn user_text(text: impl Into<String>) -> Self {
Self::new(MessageRole::User, vec![Part::text(text)])
}
pub fn with_id(mut self, id: impl Into<String>) -> Self {
self.message_id = Some(id.into());
self
}
pub fn with_task_id(mut self, task_id: impl Into<String>) -> Self {
self.task_id = Some(task_id.into());
self
}
pub fn with_context_id(mut self, context_id: impl Into<String>) -> Self {
self.context_id = Some(context_id.into());
self
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "camelCase")]
pub enum Part {
#[serde(rename = "text")]
Text {
text: String,
},
#[serde(rename = "file")]
File {
file: FileContent,
},
#[serde(rename = "data")]
Data {
data: serde_json::Value,
},
#[serde(other)]
Unknown,
}
impl Part {
pub fn text(text: impl Into<String>) -> Self {
Part::Text { text: text.into() }
}
pub fn file_uri(uri: impl Into<String>, mime_type: Option<String>) -> Self {
Part::File {
file: FileContent::Uri { uri: uri.into(), mime_type },
}
}
pub fn file_bytes(bytes: Vec<u8>, mime_type: Option<String>, name: Option<String>) -> Self {
Part::File {
file: FileContent::Bytes {
bytes: base64::Engine::encode(&base64::engine::general_purpose::STANDARD, &bytes),
mime_type,
name,
},
}
}
pub fn data(data: serde_json::Value) -> Self {
Part::Data { data }
}
pub fn as_text(&self) -> Option<&str> {
match self {
Part::Text { text } => Some(text),
_ => None,
}
}
pub fn is_unknown(&self) -> bool {
matches!(self, Part::Unknown)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum FileContent {
Uri {
uri: String,
#[serde(skip_serializing_if = "Option::is_none")]
mime_type: Option<String>,
},
Bytes {
bytes: String,
#[serde(skip_serializing_if = "Option::is_none")]
mime_type: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Artifact {
pub id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
pub parts: Vec<Part>,
#[serde(skip_serializing_if = "Option::is_none")]
pub index: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, serde_json::Value>>,
}
impl Artifact {
pub fn text(id: impl Into<String>, text: impl Into<String>) -> Self {
Self {
id: id.into(),
name: None,
description: None,
parts: vec![Part::text(text)],
index: None,
metadata: None,
}
}
pub fn file(id: impl Into<String>, file: FileContent) -> Self {
Self {
id: id.into(),
name: None,
description: None,
parts: vec![Part::File { file }],
index: None,
metadata: None,
}
}
pub fn with_name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
pub fn with_description(mut self, description: impl Into<String>) -> Self {
self.description = Some(description.into());
self
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Task {
pub id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub context_id: Option<String>,
pub status: TaskStatus,
#[serde(skip_serializing_if = "Vec::is_empty", default)]
pub artifacts: Vec<Artifact>,
#[serde(skip_serializing_if = "Vec::is_empty", default)]
pub history: Vec<Message>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, serde_json::Value>>,
#[serde(default = "default_task_kind")]
pub kind: String,
}
fn default_task_kind() -> String {
"task".to_string()
}
impl Task {
pub fn new() -> Self {
Self {
id: uuid::Uuid::new_v4().to_string(),
context_id: None,
status: TaskStatus::default(),
artifacts: Vec::new(),
history: Vec::new(),
metadata: None,
kind: "task".to_string(),
}
}
pub fn with_id(id: impl Into<String>) -> Self {
Self {
id: id.into(),
context_id: None,
status: TaskStatus::default(),
artifacts: Vec::new(),
history: Vec::new(),
metadata: None,
kind: "task".to_string(),
}
}
pub fn with_context_id(mut self, context_id: impl Into<String>) -> Self {
self.context_id = Some(context_id.into());
self
}
pub fn state(&self) -> TaskState {
self.status.state
}
pub fn is_terminal(&self) -> bool {
self.status.state.is_terminal()
}
pub fn is_cancelable(&self) -> bool {
self.status.state.is_cancelable()
}
pub fn update_status(&mut self, state: TaskState, message: Option<Message>) {
self.status = match message {
Some(msg) => TaskStatus::with_message(state, msg),
None => TaskStatus::new(state),
};
}
pub fn add_artifact(&mut self, artifact: Artifact) {
self.artifacts.push(artifact);
}
pub fn add_message(&mut self, message: Message) {
self.history.push(message);
}
}
impl Default for Task {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_task_state_terminal() {
assert!(!TaskState::Submitted.is_terminal());
assert!(!TaskState::Working.is_terminal());
assert!(TaskState::Completed.is_terminal());
assert!(TaskState::Failed.is_terminal());
assert!(TaskState::Canceled.is_terminal());
}
#[test]
fn test_task_state_cancelable() {
assert!(TaskState::Submitted.is_cancelable());
assert!(TaskState::Working.is_cancelable());
assert!(TaskState::InputRequired.is_cancelable());
assert!(!TaskState::Completed.is_cancelable());
assert!(!TaskState::Failed.is_cancelable());
}
#[test]
fn test_message_creation() {
let msg = Message::agent_text("Hello, world!");
assert_eq!(msg.role, MessageRole::Agent);
assert_eq!(msg.parts.len(), 1);
assert_eq!(msg.parts[0].as_text(), Some("Hello, world!"));
}
#[test]
fn test_task_lifecycle() {
let mut task = Task::new();
assert_eq!(task.state(), TaskState::Submitted);
assert!(!task.is_terminal());
assert!(task.is_cancelable());
task.update_status(TaskState::Working, None);
assert_eq!(task.state(), TaskState::Working);
task.update_status(TaskState::Completed, Some(Message::agent_text("Task completed")));
assert!(task.is_terminal());
assert!(!task.is_cancelable());
}
#[test]
fn test_part_serialization() {
let text_part = Part::text("Hello");
let json = serde_json::to_string(&text_part).expect("serialize");
assert!(json.contains("\"type\":\"text\""));
let data_part = Part::data(serde_json::json!({"key": "value"}));
let json = serde_json::to_string(&data_part).expect("serialize");
assert!(json.contains("\"type\":\"data\""));
}
#[test]
fn test_artifact_creation() {
let artifact = Artifact::text("art-1", "Generated content")
.with_name("output.txt")
.with_description("The generated output file");
assert_eq!(artifact.id, "art-1");
assert_eq!(artifact.name, Some("output.txt".to_string()));
assert_eq!(artifact.parts.len(), 1);
}
}