diff --git a/README.md b/README.md index c3906ce..6e51fa1 100644 --- a/README.md +++ b/README.md @@ -339,6 +339,54 @@ field works without upgrading the CLI. `--from-version latest|N` starts from an existing version and applies your overrides on top, replacing whole top-level keys rather than deep-merging. +### Talking to an agent (A2A) + +Every agent bound to a workspace answers over the +[A2A protocol](https://a2a-protocol.org). The CLI speaks A2A v1.0 over +HTTP+JSON; `--workspace` defaults to the remembered one, as for `agent bind`. + +```bash +memorylake agent card [--workspace ] + +memorylake agent send --text "..." [--text "..."] \ + [--context CTX] [--task TASK] [--stream | --no-wait] [--raw] \ + [--actor ID] [--project ID] [--read-only-project ID]... [--skip-memory] \ + [--metadata-json '{"overrides":{...}}'] [--workspace ] +memorylake agent send --message-json '[{"text":"..."},{"url":"...","mediaType":"image/png"}]' +memorylake agent send --message-file parts.json + +memorylake agent task list [--context CTX] [--status STATE] \ + [--page-size N] [--page-token TOK] [--after TIMESTAMP] [--history-length N] [--artifacts] +memorylake agent task get [--history-length N] +memorylake agent task cancel +memorylake agent task feedback --rating up|down [--comment TEXT] +``` + +`send` waits for the answer and prints the reply text on stdout; the task id, +context id and final state go to stderr, so a script can pipe the reply and +still know how to continue: + +``` +$ memorylake agent send agent-… --text "Summarize yesterday's standup" +The team agreed to … +task run-… context 5fdb… state TASK_STATE_COMPLETED +``` + +Pass `--context` to keep talking in the same thread. A task that ends in +`TASK_STATE_INPUT_REQUIRED` needs more from you: reply with `--task +--context `. `--stream` prints the reply as it is produced; `--no-wait` +returns the task as soon as it exists (poll it with `agent task get`); `--raw` +prints the protocol response as JSON instead of the reply text. + +The `--actor`, `--project`, `--read-only-project` and `--skip-memory` flags set +MemoryLake's extension of the request (`metadata.memorylake`): whose message it +is, which projects the agent may read and write, and whether the exchange is +remembered at all. Anything else the extension accepts, such as `overrides`, +goes through `--metadata-json`. + +Feedback is a MemoryLake extension to A2A. `task feedback` records a rating on +the task; `task get` reads it back under `metadata."task-feedback/v1"`. + ### Search ```bash diff --git a/crates/cli/src/commands/agent.rs b/crates/cli/src/commands/agent.rs index 933cde3..b232d42 100644 --- a/crates/cli/src/commands/agent.rs +++ b/crates/cli/src/commands/agent.rs @@ -2,8 +2,10 @@ //! //! Agent *identity* (name, description, metadata) changes in place via //! `agent update`. Agent *configuration* (model, policies, prompt, …) is -//! immutable and changes only by creating a new version. +//! immutable and changes only by creating a new version. Talking to a bound +//! agent (`card`, `send`, `task`) goes over A2A and lives in [`a2a`]. +mod a2a; mod body; use anyhow::{Context, Result}; @@ -17,6 +19,7 @@ use memorylake_core::{Client, Paths, ResolveOverrides, resolve}; use std::path::PathBuf; use super::require_workspace; +use a2a::{SendArgs, TaskCommand, run_card, run_send, run_task, task_workspace_flag}; use body::{FromVersion, load_config_body, reject_config_fields, require_field, set_scalar}; /// Agent subcommands. @@ -138,6 +141,26 @@ pub enum AgentCommand { #[arg(long = "name")] name_fuzzy: Option, }, + /// Show an agent's A2A card: capabilities, protocol bindings, extensions. + Card { + /// Agent id (must be bound to the workspace). + agent_id: String, + /// Workspace the agent is bound in. + /// + /// Defaults to the workspace remembered by `workspace use`. + #[arg(long)] + workspace: Option, + }, + /// Send a message to an agent and print its reply. + /// + /// Waits for the answer by default. The reply text goes to stdout; the + /// task and context ids needed to continue go to stderr. + Send(SendArgs), + /// Inspect, cancel and rate the tasks an agent has run. + Task { + #[command(subcommand)] + command: TaskCommand, + }, } /// `agent version` subcommands. @@ -303,6 +326,22 @@ pub fn run(command: AgentCommand, profile: Option, base_url: Option { + let workspace = require_workspace(&paths, &runtime.profile, workspace)?; + run_card(&client, &workspace, &agent_id)?; + } + AgentCommand::Send(args) => { + let workspace = require_workspace(&paths, &runtime.profile, args.workspace.clone())?; + run_send(&client, &workspace, args)?; + } + AgentCommand::Task { command } => { + let workspace = + require_workspace(&paths, &runtime.profile, task_workspace_flag(&command))?; + run_task(&client, &workspace, command)?; + } } Ok(()) diff --git a/crates/cli/src/commands/agent/a2a.rs b/crates/cli/src/commands/agent/a2a.rs new file mode 100644 index 0000000..b1b33d8 --- /dev/null +++ b/crates/cli/src/commands/agent/a2a.rs @@ -0,0 +1,773 @@ +//! `memorylake agent card|send|task`: talking to a bound agent over A2A. +//! +//! `send` is the conversation itself; `task` is the bookkeeping around it. Both +//! address the agent inside a workspace, so they take `--workspace` with the +//! same default as `agent bind`. +//! +//! Output is split so a script can pipe it: the agent's words go to stdout, +//! the ids needed to continue (task, context) go to stderr. `--raw` turns +//! that off and prints the protocol response as JSON. + +use std::io::Write; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{SystemTime, UNIX_EPOCH}; + +use anyhow::{Context, Result, bail}; +use clap::{Args, Subcommand, ValueEnum}; +use memorylake_core::Client; +use memorylake_core::api::agents::a2a::{ + FEEDBACK_COMMENT_MAX_CHARS, ListTasksParams, MemorylakeExtension, Message, ROLE_USER, Rating, + SendConfiguration, SendMessageRequest, SendMetadata, TASK_STATE_INPUT_REQUIRED, + TaskFeedbackRequest, cancel_task, get_agent_card, get_task, list_tasks, send_message, + stream_message, submit_task_feedback, text_part, +}; +use serde_json::{Map, Value}; + +use super::super::print_json; + +/// Arguments of `agent send`. +#[derive(Debug, Args)] +pub struct SendArgs { + /// Agent to talk to (must be bound to the workspace). + pub agent_id: String, + /// Message text. Repeat for several text parts. + #[arg(long, value_name = "TEXT")] + pub text: Vec, + /// Full A2A message parts as an inline JSON array, for non-text parts. + #[arg(long, value_name = "JSON")] + pub message_json: Option, + /// Full A2A message parts read from a JSON file. + #[arg(long, value_name = "PATH")] + pub message_file: Option, + /// Continue an existing conversation thread. + #[arg(long, value_name = "CONTEXT_ID")] + pub context: Option, + /// Reply to a task that stopped in `TASK_STATE_INPUT_REQUIRED`. + #[arg(long, value_name = "TASK_ID")] + pub task: Option, + /// Print the agent's reply as it is produced. + #[arg(long, conflicts_with = "no_wait")] + pub stream: bool, + /// Return as soon as the task exists instead of waiting for the answer. + /// + /// Prints the task as JSON; follow it with `agent task get`. + #[arg(long)] + pub no_wait: bool, + /// Print the protocol response as JSON instead of the reply text. + #[arg(long)] + pub raw: bool, + /// Actor the message is attributed to (`metadata.memorylake.actorId`). + #[arg(long, value_name = "ACTOR_ID")] + pub actor: Option, + /// Project the agent reads from and writes memories to. + #[arg(long, value_name = "PROJECT_ID")] + pub project: Option, + /// Project the agent may read but not write. Repeatable. + #[arg(long = "read-only-project", value_name = "PROJECT_ID")] + pub read_only_projects: Vec, + /// Do not extract memories from this exchange. + #[arg(long)] + pub skip_memory: bool, + /// Extra `metadata.memorylake` keys as a JSON object (e.g. `overrides`). + /// + /// Merged under the flags above; a key set both ways is rejected. + #[arg(long, value_name = "JSON")] + pub metadata_json: Option, + /// Workspace the agent is bound in. + /// + /// Defaults to the workspace remembered by `workspace use`. + #[arg(long)] + pub workspace: Option, +} + +/// `agent task` subcommands. +#[derive(Debug, Subcommand)] +pub enum TaskCommand { + /// List an agent's tasks, newest first. + List { + /// Agent whose tasks to list. + agent_id: String, + /// Only tasks in this conversation thread. + #[arg(long, value_name = "CONTEXT_ID")] + context: Option, + /// Only tasks in this state (e.g. `TASK_STATE_WORKING`). + #[arg(long, value_name = "STATE")] + status: Option, + /// Tasks per page (1–100). + #[arg(long, value_parser = clap::value_parser!(u32).range(1..=100))] + page_size: Option, + /// `nextPageToken` from the previous page. + #[arg(long)] + page_token: Option, + /// Only tasks whose status changed after this ISO 8601 instant. + #[arg(long, value_name = "TIMESTAMP")] + after: Option, + /// History messages to include per task. + #[arg(long)] + history_length: Option, + /// Include each task's artifacts. + #[arg(long)] + artifacts: bool, + /// Workspace the agent is bound in. + /// + /// Defaults to the workspace remembered by `workspace use`. + #[arg(long)] + workspace: Option, + }, + /// Get one task. + Get { + /// Agent that owns the task. + agent_id: String, + /// Task id. + task_id: String, + /// History messages to include. + #[arg(long)] + history_length: Option, + /// Workspace the agent is bound in. + /// + /// Defaults to the workspace remembered by `workspace use`. + #[arg(long)] + workspace: Option, + }, + /// Cancel a running task. + Cancel { + /// Agent that owns the task. + agent_id: String, + /// Task id. + task_id: String, + /// Workspace the agent is bound in. + /// + /// Defaults to the workspace remembered by `workspace use`. + #[arg(long)] + workspace: Option, + }, + /// Rate a task's result. + Feedback { + /// Agent that owns the task. + agent_id: String, + /// Task id. + task_id: String, + /// Thumbs up or down. + #[arg(long, value_enum)] + rating: RatingArg, + /// Free-text comment (at most 2000 characters). + #[arg(long)] + comment: Option, + /// Workspace the agent is bound in. + /// + /// Defaults to the workspace remembered by `workspace use`. + #[arg(long)] + workspace: Option, + }, +} + +/// `--rating` as spelled on the command line. +#[derive(Debug, Clone, Copy, PartialEq, Eq, ValueEnum)] +pub enum RatingArg { + /// The result was good. + Up, + /// The result was bad. + Down, +} + +impl From for Rating { + fn from(rating: RatingArg) -> Self { + match rating { + RatingArg::Up => Self::Up, + RatingArg::Down => Self::Down, + } + } +} + +/// `agent card`. +pub fn run_card(client: &Client, workspace: &str, agent_id: &str) -> Result<()> { + let card = get_agent_card(client, workspace, agent_id) + .with_context(|| format!("get agent card of `{agent_id}` in workspace `{workspace}`"))?; + print_json(&card) +} + +/// `agent send`. +pub fn run_send(client: &Client, workspace: &str, args: SendArgs) -> Result<()> { + let request = build_request(&args)?; + let agent_id = &args.agent_id; + + if args.stream { + let events = stream_message(client, workspace, agent_id, &request) + .with_context(|| format!("stream message to agent `{agent_id}`"))?; + return print_stream(events, args.raw); + } + + let response = send_message(client, workspace, agent_id, &request) + .with_context(|| format!("send message to agent `{agent_id}`"))?; + if args.raw || args.no_wait { + return print_json(&response); + } + print_reply(&response) +} + +/// `agent task ...`. +pub fn run_task(client: &Client, workspace: &str, command: TaskCommand) -> Result<()> { + match command { + TaskCommand::List { + agent_id, + context, + status, + page_size, + page_token, + after, + history_length, + artifacts, + workspace: _, + } => { + let params = ListTasksParams { + context_id: context, + status, + page_size, + page_token, + history_length, + status_timestamp_after: after, + include_artifacts: artifacts.then_some(true), + }; + let data = list_tasks(client, workspace, &agent_id, ¶ms) + .with_context(|| format!("list tasks of agent `{agent_id}`"))?; + print_json(&data) + } + TaskCommand::Get { + agent_id, + task_id, + history_length, + workspace: _, + } => { + let data = get_task(client, workspace, &agent_id, &task_id, history_length) + .with_context(|| format!("get task `{task_id}` of agent `{agent_id}`"))?; + print_json(&data) + } + TaskCommand::Cancel { + agent_id, + task_id, + workspace: _, + } => { + let data = cancel_task(client, workspace, &agent_id, &task_id) + .with_context(|| format!("cancel task `{task_id}` of agent `{agent_id}`"))?; + print_json(&data) + } + TaskCommand::Feedback { + agent_id, + task_id, + rating, + comment, + workspace: _, + } => { + if let Some(comment) = &comment { + let chars = comment.chars().count(); + if chars > FEEDBACK_COMMENT_MAX_CHARS { + bail!( + "--comment is {chars} characters; the limit is {FEEDBACK_COMMENT_MAX_CHARS}" + ); + } + } + let request = TaskFeedbackRequest { + rating: rating.into(), + comment, + }; + let data = submit_task_feedback(client, workspace, &agent_id, &task_id, &request) + .with_context(|| format!("rate task `{task_id}` of agent `{agent_id}`"))?; + print_json(&data) + } + } +} + +/// The `--workspace` flag of a `task` subcommand, for the shared resolver. +pub fn task_workspace_flag(command: &TaskCommand) -> Option { + match command { + TaskCommand::List { workspace, .. } + | TaskCommand::Get { workspace, .. } + | TaskCommand::Cancel { workspace, .. } + | TaskCommand::Feedback { workspace, .. } => workspace.clone(), + } +} + +/// Turn `send` flags into the request body. +pub fn build_request(args: &SendArgs) -> Result { + let parts = build_parts( + &args.text, + args.message_json.as_deref(), + args.message_file.as_deref(), + )?; + + let configuration = SendConfiguration { + return_immediately: args.no_wait.then_some(true), + history_length: None, + }; + + let mut memorylake = MemorylakeExtension { + actor_id: args.actor.clone(), + read_write_project_id: args.project.clone(), + read_only_project_ids: args.read_only_projects.clone(), + skip_memory: args.skip_memory.then_some(true), + extra: Map::new(), + }; + if let Some(json) = &args.metadata_json { + memorylake.extra = parse_extra_metadata(json, &memorylake)?; + } + + Ok(SendMessageRequest { + message: Message { + role: ROLE_USER.into(), + message_id: new_message_id(), + parts, + context_id: args.context.clone(), + task_id: args.task.clone(), + }, + configuration: (!configuration.is_empty()).then_some(configuration), + metadata: (!memorylake.is_empty()).then_some(SendMetadata { memorylake }), + }) +} + +/// A message id unique enough for the server to tell messages apart. +/// +/// A2A only needs the id to be unique within its context; time, process and a +/// counter cover that without pulling in a UUID dependency. +fn new_message_id() -> String { + static NEXT: AtomicU64 = AtomicU64::new(0); + let nanos = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_nanos()) + .unwrap_or_default(); + format!( + "cli-{nanos:x}-{:x}-{}", + std::process::id(), + NEXT.fetch_add(1, Ordering::Relaxed) + ) +} + +/// Build message parts from the mutually exclusive content flags. +/// +/// Same rule as `conversation message append`: `--text` for the common case, +/// a JSON source for anything else, never both. +fn build_parts( + texts: &[String], + message_json: Option<&str>, + message_file: Option<&Path>, +) -> Result> { + let json_source = match (message_json, message_file) { + (Some(_), Some(_)) => bail!("pass --message-json or --message-file, not both"), + (Some(inline), None) => Some((inline.to_string(), "--message-json".to_string())), + (None, Some(path)) => { + let text = std::fs::read_to_string(path) + .with_context(|| format!("read message file {}", path.display()))?; + Some((text, format!("message file {}", path.display()))) + } + (None, None) => None, + }; + + match (texts.is_empty(), json_source) { + (false, Some((_, source))) => { + bail!("--text and {source} both set the message; pass only one") + } + (true, None) => bail!( + "a message is required: pass --text , \ + or --message-json / --message-file for non-text parts" + ), + (false, None) => Ok(texts.iter().map(text_part).collect()), + (true, Some((json, source))) => parse_parts(&json, &source), + } +} + +/// Parse a JSON array of A2A parts. Each must be an object; what is inside is +/// the server's business. +fn parse_parts(json: &str, source: &str) -> Result> { + let value: Value = + serde_json::from_str(json).with_context(|| format!("parse JSON from {source}"))?; + let Value::Array(items) = value else { + bail!( + "{source} must hold a JSON array of message parts\n\ + example: [{{\"text\": \"hello\"}}]" + ); + }; + if items.is_empty() { + bail!("{source} holds an empty array; a message needs at least one part"); + } + for (index, item) in items.iter().enumerate() { + if !item.is_object() { + bail!("{source}: part {index} must be a JSON object"); + } + } + Ok(items) +} + +/// Parse `--metadata-json` and refuse keys the dedicated flags already own. +fn parse_extra_metadata(json: &str, flags: &MemorylakeExtension) -> Result> { + let value: Value = + serde_json::from_str(json).with_context(|| "parse JSON from --metadata-json")?; + let Value::Object(extra) = value else { + bail!( + "--metadata-json must hold a JSON object, e.g. {{\"overrides\": {{\"model\": \"...\"}}}}" + ); + }; + let owned = [ + ("actorId", flags.actor_id.is_some(), "--actor"), + ( + "readWriteProjectId", + flags.read_write_project_id.is_some(), + "--project", + ), + ( + "readOnlyProjectIds", + !flags.read_only_project_ids.is_empty(), + "--read-only-project", + ), + ("skipMemory", flags.skip_memory.is_some(), "--skip-memory"), + ]; + for (key, set_by_flag, flag) in owned { + if extra.contains_key(key) { + if set_by_flag { + bail!("`{key}` is set both by {flag} and in --metadata-json; pass only one"); + } + bail!("set `{key}` with {flag} rather than in --metadata-json"); + } + } + Ok(extra) +} + +/// Print a blocking `send` response: the reply text, then where to continue. +fn print_reply(response: &Value) -> Result<()> { + let task = response.get("task").unwrap_or(response); + let text = reply_text(task); + + if text.is_empty() { + // Nothing the CLI knows how to read as words; show everything rather + // than nothing. + print_json(response)?; + } else { + println!("{text}"); + } + print_continuation(task); + Ok(()) +} + +/// The agent's words in a finished task. +/// +/// Artifacts hold the result of a completed task; the status message holds +/// what the agent said when it stopped for input. A bare `message` response +/// (no task) is read the same way. +fn reply_text(task: &Value) -> String { + let mut text = String::new(); + if let Some(artifacts) = task.get("artifacts").and_then(Value::as_array) { + for artifact in artifacts { + append_parts_text(&mut text, artifact.get("parts")); + } + } + if text.is_empty() { + append_parts_text(&mut text, task.pointer("/status/message/parts")); + } + if text.is_empty() { + // `{"message": {...}}` responses, or a task-less reply. + append_parts_text(&mut text, task.get("parts")); + } + text +} + +fn append_parts_text(out: &mut String, parts: Option<&Value>) { + let Some(parts) = parts.and_then(Value::as_array) else { + return; + }; + for part in parts { + if let Some(fragment) = part.get("text").and_then(Value::as_str) { + out.push_str(fragment); + } + } +} + +/// Tell the caller how to continue, on stderr so stdout stays the reply. +fn print_continuation(task: &Value) { + let id = task.get("id").and_then(Value::as_str).unwrap_or("?"); + let context = task.get("contextId").and_then(Value::as_str).unwrap_or("?"); + let state = task + .pointer("/status/state") + .and_then(Value::as_str) + .unwrap_or("?"); + eprintln!("task {id} context {context} state {state}"); + if state == TASK_STATE_INPUT_REQUIRED { + eprintln!("the agent needs more input; reply with --task {id} --context {context}"); + } +} + +/// Print a `message:stream` as it arrives. +/// +/// Text mode writes each fragment the moment it lands and flushes, so a +/// caller watching the terminal sees the reply grow. Raw mode prints one +/// compact JSON document per event. +/// +/// Production streams the reply twice: token by token in `statusUpdate` +/// events, then once more in full as the closing `artifactUpdate`. The +/// artifact is only printed when no status text arrived, so an agent that +/// streams nothing still shows its result and one that streams does not +/// repeat itself. +fn print_stream(events: I, raw: bool) -> Result<()> +where + I: Iterator>, +{ + let stdout = std::io::stdout(); + let mut out = stdout.lock(); + let mut last_task: Option = None; + let mut streamed_status_text = false; + let mut artifact_text = String::new(); + + for event in events { + let event = event.context("read agent reply stream")?; + if raw { + writeln!(out, "{event}")?; + out.flush()?; + continue; + } + + let fragment = status_text(&event); + if !fragment.is_empty() { + out.write_all(fragment.as_bytes())?; + out.flush()?; + streamed_status_text = true; + } + append_parts_text( + &mut artifact_text, + event.pointer("/artifactUpdate/artifact/parts"), + ); + if let Some(task) = event_task_summary(&event) { + last_task = Some(task); + } + } + + if raw { + return Ok(()); + } + if streamed_status_text { + writeln!(out)?; + } else if !artifact_text.is_empty() { + writeln!(out, "{artifact_text}")?; + } + if let Some(task) = last_task { + print_continuation(&task); + } + Ok(()) +} + +/// The words a stream event carries as the agent speaks them: the status +/// message of a `statusUpdate`, or a bare `message` response. +fn status_text(event: &Value) -> String { + let mut text = String::new(); + append_parts_text( + &mut text, + event.pointer("/statusUpdate/status/message/parts"), + ); + append_parts_text(&mut text, event.pointer("/message/parts")); + text +} + +/// The task id, context and state named by one stream event, in the shape +/// [`print_continuation`] reads. +fn event_task_summary(event: &Value) -> Option { + if let Some(task) = event.get("task") { + return Some(task.clone()); + } + let update = event + .get("statusUpdate") + .or_else(|| event.get("artifactUpdate"))?; + let mut summary = Map::new(); + if let Some(id) = update.get("taskId") { + summary.insert("id".into(), id.clone()); + } + if let Some(context) = update.get("contextId") { + summary.insert("contextId".into(), context.clone()); + } + if let Some(status) = update.get("status") { + summary.insert("status".into(), status.clone()); + } + Some(Value::Object(summary)) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn args(agent: &str) -> SendArgs { + SendArgs { + agent_id: agent.into(), + text: vec![], + message_json: None, + message_file: None, + context: None, + task: None, + stream: false, + no_wait: false, + raw: false, + actor: None, + project: None, + read_only_projects: vec![], + skip_memory: false, + metadata_json: None, + workspace: None, + } + } + + #[test] + fn text_flags_become_text_parts_and_nothing_else_is_sent() { + let mut a = args("agt-1"); + a.text = vec!["hello".into(), "world".into()]; + let request = build_request(&a).unwrap(); + let body = serde_json::to_value(&request).unwrap(); + assert_eq!(body["message"]["role"], "ROLE_USER"); + assert_eq!( + body["message"]["parts"], + json!([{"text": "hello"}, {"text": "world"}]) + ); + assert!( + body["message"]["messageId"] + .as_str() + .unwrap() + .starts_with("cli-") + ); + assert!(body.get("configuration").is_none(), "{body}"); + assert!(body.get("metadata").is_none(), "{body}"); + } + + #[test] + fn message_ids_differ_between_calls() { + assert_ne!(new_message_id(), new_message_id()); + } + + #[test] + fn scope_flags_land_under_metadata_memorylake() { + let mut a = args("agt-1"); + a.text = vec!["hi".into()]; + a.actor = Some("act-1".into()); + a.project = Some("proj-1".into()); + a.read_only_projects = vec!["proj-2".into(), "proj-3".into()]; + a.skip_memory = true; + a.context = Some("ctx".into()); + a.task = Some("run-1".into()); + a.no_wait = true; + let body = serde_json::to_value(build_request(&a).unwrap()).unwrap(); + assert_eq!( + body["metadata"]["memorylake"], + json!({ + "actorId": "act-1", + "readWriteProjectId": "proj-1", + "readOnlyProjectIds": ["proj-2", "proj-3"], + "skipMemory": true + }) + ); + assert_eq!(body["configuration"], json!({"returnImmediately": true})); + assert_eq!(body["message"]["contextId"], "ctx"); + assert_eq!(body["message"]["taskId"], "run-1"); + } + + #[test] + fn metadata_json_is_merged_but_may_not_shadow_a_flag() { + let mut a = args("agt-1"); + a.text = vec!["hi".into()]; + a.metadata_json = Some(r#"{"overrides":{"model":"x"}}"#.into()); + let body = serde_json::to_value(build_request(&a).unwrap()).unwrap(); + assert_eq!( + body["metadata"]["memorylake"], + json!({"overrides": {"model": "x"}}) + ); + + a.metadata_json = Some(r#"{"skipMemory":true}"#.into()); + let err = build_request(&a).unwrap_err().to_string(); + assert!(err.contains("--skip-memory"), "{err}"); + + a.skip_memory = true; + let err = build_request(&a).unwrap_err().to_string(); + assert!(err.contains("both"), "{err}"); + + a.metadata_json = Some("[]".into()); + let err = build_request(&a).unwrap_err().to_string(); + assert!(err.contains("JSON object"), "{err}"); + } + + #[test] + fn message_content_is_required_and_exclusive() { + let a = args("agt-1"); + let err = build_request(&a).unwrap_err().to_string(); + assert!(err.contains("--text"), "{err}"); + + let mut a = args("agt-1"); + a.text = vec!["hi".into()]; + a.message_json = Some(r#"[{"text":"x"}]"#.into()); + let err = build_request(&a).unwrap_err().to_string(); + assert!(err.contains("pass only one"), "{err}"); + + let mut a = args("agt-1"); + a.message_json = + Some(r#"[{"text":"x"},{"url":"https://e/x.png","mediaType":"image/png"}]"#.into()); + let body = serde_json::to_value(build_request(&a).unwrap()).unwrap(); + assert_eq!(body["message"]["parts"].as_array().unwrap().len(), 2); + + a.message_json = Some(r#"{"text":"x"}"#.into()); + let err = build_request(&a).unwrap_err().to_string(); + assert!(err.contains("JSON array"), "{err}"); + + a.message_json = Some("[]".into()); + let err = build_request(&a).unwrap_err().to_string(); + assert!(err.contains("empty array"), "{err}"); + + a.message_json = Some(r#"["x"]"#.into()); + let err = build_request(&a).unwrap_err().to_string(); + assert!(err.contains("part 0"), "{err}"); + } + + #[test] + fn reply_text_prefers_artifacts_then_the_status_message() { + let completed = json!({ + "id": "run-1", + "status": {"state": "TASK_STATE_COMPLETED"}, + "artifacts": [{"parts": [{"text": "PO"}, {"text": "NG"}]}], + "history": [{"role": "ROLE_AGENT", "parts": [{"text": "ignored"}]}] + }); + assert_eq!(reply_text(&completed), "PONG"); + + let input_required = json!({ + "id": "run-2", + "status": { + "state": "TASK_STATE_INPUT_REQUIRED", + "message": {"parts": [{"text": "Which file?"}]} + }, + "artifacts": [] + }); + assert_eq!(reply_text(&input_required), "Which file?"); + + let bare_message = json!({"role": "ROLE_AGENT", "parts": [{"text": "hi"}]}); + assert_eq!(reply_text(&bare_message), "hi"); + + assert_eq!(reply_text(&json!({"id": "run-3"})), ""); + } + + #[test] + fn stream_events_yield_their_text_and_task_summary() { + let first = json!({"task": {"id": "run-1", "contextId": "ctx", "status": {"state": "TASK_STATE_WORKING"}}}); + assert_eq!(status_text(&first), ""); + assert_eq!(event_task_summary(&first).unwrap()["id"], "run-1"); + + let update = json!({"statusUpdate": { + "taskId": "run-1", "contextId": "ctx", + "status": {"state": "TASK_STATE_WORKING", "message": {"parts": [{"text": "1"}, {"text": "\n"}]}} + }}); + assert_eq!(status_text(&update), "1\n"); + let summary = event_task_summary(&update).unwrap(); + assert_eq!(summary["id"], "run-1"); + assert_eq!(summary["contextId"], "ctx"); + assert_eq!(summary["status"]["state"], "TASK_STATE_WORKING"); + + let artifact = json!({"artifactUpdate": {"taskId": "run-1", "artifact": {"parts": [{"text": "done"}]}}}); + assert_eq!( + status_text(&artifact), + "", + "artifact text is the consolidated reply, not a streamed fragment" + ); + assert_eq!(event_task_summary(&artifact).unwrap()["id"], "run-1"); + + assert!(event_task_summary(&json!({"unknown": {}})).is_none()); + } +} diff --git a/crates/cli/tests/agent/live.rs b/crates/cli/tests/agent/live.rs index 54d525f..1c14988 100644 --- a/crates/cli/tests/agent/live.rs +++ b/crates/cli/tests/agent/live.rs @@ -332,3 +332,189 @@ fn bound_agent_ids(bindings: &Value) -> Vec { }) .unwrap_or_default() } + +/// The account's default workspace and the built-in agent bound to it. +/// +/// Every account has both, so the A2A tests need no fixture of their own — +/// and creating one would not do: a freshly created agent has no runtime to +/// answer over A2A until it is configured with a model. +fn default_workspace_and_agent(home: &Path) -> (String, String) { + // The CI account carries hundreds of scratch workspaces from earlier runs, + // so the default one may sit several pages in. + let mut token: Option = None; + let workspace = loop { + let mut ws_args = vec!["ws", "list", "--page-size", "100"]; + if let Some(token) = &token { + ws_args.extend(["--continuation-token", token.as_str()]); + } + let page = parse_json(&assert_success(&run(home, &ws_args), &ws_args), "ws list"); + if let Some(found) = page["items"] + .as_array() + .expect("items") + .iter() + .find(|ws| ws["custom_id"] == "_sys_default_workspace") + { + break str_field(found, "id", "ws list").to_string(); + } + match page["continuation_token"].as_str() { + Some(next) if !next.is_empty() => token = Some(next.to_string()), + _ => panic!("account has no `_sys_default_workspace` on any page"), + } + }; + + let bound_args = ["agent", "bindings", "--workspace", workspace.as_str()]; + let bound = parse_json( + &assert_success(&run(home, &bound_args), &bound_args), + "bindings", + ); + let agent = bound["items"] + .as_array() + .expect("items") + .iter() + .find(|b| b["custom_id"] == "_sys_builtin_default_agent") + .map(|b| str_field(b, "agent_id", "bindings").to_string()) + .expect("default agent is bound to the default workspace"); + + (workspace, agent) +} + +#[test] +fn a2a_round_trip_against_the_default_agent() { + let api_key = require_api_key(); + let home = temp_home(); + login_default(&home, &api_key); + let (workspace, agent) = default_workspace_and_agent(&home); + let ws = workspace.as_str(); + let agent = agent.as_str(); + + // 1. The card names the v1.0 HTTP+JSON binding this CLI speaks. + let card_args = ["agent", "card", agent, "--workspace", ws]; + let card = parse_json(&assert_success(&run(&home, &card_args), &card_args), "card"); + let interfaces = card["supportedInterfaces"].as_array().expect("interfaces"); + assert!( + interfaces.iter().any(|i| { + i["protocolBinding"] == "HTTP+JSON" + && i["protocolVersion"] == "1.0" + && i["url"].as_str().is_some_and(|u| u.ends_with("/a2a")) + }), + "card advertises no v1.0 HTTP+JSON binding on the unversioned path: {card}" + ); + + // 2. A blocking send prints the reply and names the task on stderr. + // `--skip-memory` keeps the probe out of the account's memories. + let send_args = [ + "agent", + "send", + agent, + "--workspace", + ws, + "--text", + "Reply with exactly the word PONG and nothing else.", + "--skip-memory", + ]; + let output = run(&home, &send_args); + let stdout = assert_success(&output, &send_args); + assert!( + stdout.to_uppercase().contains("PONG"), + "reply should contain PONG: {stdout}" + ); + let stderr = String::from_utf8_lossy(&output.stderr); + let task_id = stderr + .lines() + .find_map(|line| line.strip_prefix("task ")) + .and_then(|rest| rest.split_whitespace().next()) + .unwrap_or_else(|| panic!("stderr names no task: {stderr}")) + .to_string(); + assert!(stderr.contains("TASK_STATE_COMPLETED"), "{stderr}"); + + // 3. The task can be read back, and it appears in the listing. + let get_args = [ + "agent", + "task", + "get", + agent, + task_id.as_str(), + "--workspace", + ws, + ]; + let task = parse_json( + &assert_success(&run(&home, &get_args), &get_args), + "task get", + ); + assert_eq!(str_field(&task, "id", "task get"), task_id); + + let list_args = [ + "agent", + "task", + "list", + agent, + "--workspace", + ws, + "--page-size", + "1", + ]; + let page = parse_json( + &assert_success(&run(&home, &list_args), &list_args), + "task list", + ); + assert_eq!(page["tasks"].as_array().map(Vec::len), Some(1), "{page}"); + + // 4. Feedback lands in the task's metadata and reads back with the + // extension header `task get` sends. + let feedback_args = [ + "agent", + "task", + "feedback", + agent, + task_id.as_str(), + "--workspace", + ws, + "--rating", + "up", + "--comment", + "memorylake-cli live test", + ]; + assert_success(&run(&home, &feedback_args), &feedback_args); + let task = parse_json( + &assert_success(&run(&home, &get_args), &get_args), + "task get after feedback", + ); + assert_eq!( + task["metadata"]["task-feedback/v1"]["rating"], "up", + "rating missing after feedback: {task}" + ); + + // 5. Cancelling a finished task is refused with the A2A reason. + let cancel_args = [ + "agent", + "task", + "cancel", + agent, + task_id.as_str(), + "--workspace", + ws, + ]; + let err = assert_failure(&run(&home, &cancel_args), &cancel_args); + assert!(err.contains("TASK_NOT_CANCELABLE"), "{err}"); + + // 6. Streaming prints the reply once, not once per token plus the artifact. + let stream_args = [ + "agent", + "send", + agent, + "--workspace", + ws, + "--text", + "Reply with exactly the word PONG and nothing else.", + "--skip-memory", + "--stream", + ]; + let stdout = assert_success(&run(&home, &stream_args), &stream_args); + assert_eq!( + stdout.to_uppercase().matches("PONG").count(), + 1, + "streamed reply printed more than once: {stdout:?}" + ); + + let _ = fs::remove_dir_all(&home); +} diff --git a/crates/cli/tests/agent/mod.rs b/crates/cli/tests/agent/mod.rs index d95b489..bfe7abd 100644 --- a/crates/cli/tests/agent/mod.rs +++ b/crates/cli/tests/agent/mod.rs @@ -2,3 +2,4 @@ mod live; mod offline; +mod wire; diff --git a/crates/cli/tests/agent/offline.rs b/crates/cli/tests/agent/offline.rs index f007ee9..6aa3473 100644 --- a/crates/cli/tests/agent/offline.rs +++ b/crates/cli/tests/agent/offline.rs @@ -57,6 +57,7 @@ fn help_lists_every_agent_subcommand() { let stdout = assert_success(&run(&home, &args), &args); for subcommand in [ "list", "create", "get", "update", "delete", "version", "bind", "unbind", "bindings", + "card", "send", "task", ] { assert!( stdout.contains(subcommand), @@ -367,3 +368,175 @@ fn bind_and_bindings_require_a_workspace() { let _ = fs::remove_dir_all(&home); } + +#[test] +fn a2a_commands_require_a_workspace() { + let home = temp_home(); + seed_credentials(&home); + + for args in [ + vec!["agent", "card", "agt-1"], + vec!["agent", "send", "agt-1", "--text", "hi"], + vec!["agent", "task", "list", "agt-1"], + vec!["agent", "task", "get", "agt-1", "run-1"], + vec!["agent", "task", "cancel", "agt-1", "run-1"], + vec![ + "agent", "task", "feedback", "agt-1", "run-1", "--rating", "up", + ], + ] { + let err = assert_failure(&run(&home, &args), &args); + assert!( + err.contains("--workspace"), + "unexpected error for {args:?}: {err}" + ); + assert_no_request_attempted(&err); + } + + let _ = fs::remove_dir_all(&home); +} + +#[test] +fn send_requires_exactly_one_message_source() { + let home = temp_home(); + seed_credentials(&home); + + let args = ["agent", "send", "agt-1", "--workspace", "ws-1"]; + let err = assert_failure(&run(&home, &args), &args); + assert!(err.contains("--text"), "{err}"); + assert_no_request_attempted(&err); + + let args = [ + "agent", + "send", + "agt-1", + "--workspace", + "ws-1", + "--text", + "hi", + "--message-json", + r#"[{"text":"hi"}]"#, + ]; + let err = assert_failure(&run(&home, &args), &args); + assert!(err.contains("pass only one"), "{err}"); + assert_no_request_attempted(&err); + + let args = [ + "agent", + "send", + "agt-1", + "--workspace", + "ws-1", + "--message-json", + r#"{"text":"hi"}"#, + ]; + let err = assert_failure(&run(&home, &args), &args); + assert!(err.contains("JSON array"), "{err}"); + assert_no_request_attempted(&err); + + let _ = fs::remove_dir_all(&home); +} + +#[test] +fn send_rejects_stream_together_with_no_wait() { + let home = temp_home(); + let args = [ + "agent", + "send", + "agt-1", + "--workspace", + "ws-1", + "--text", + "hi", + "--stream", + "--no-wait", + ]; + let err = assert_failure(&run(&home, &args), &args); + assert!( + err.contains("--stream") && err.contains("--no-wait"), + "clap should name both flags: {err}" + ); + let _ = fs::remove_dir_all(&home); +} + +#[test] +fn send_rejects_metadata_json_that_duplicates_a_flag() { + let home = temp_home(); + seed_credentials(&home); + let args = [ + "agent", + "send", + "agt-1", + "--workspace", + "ws-1", + "--text", + "hi", + "--skip-memory", + "--metadata-json", + r#"{"skipMemory":false}"#, + ]; + let err = assert_failure(&run(&home, &args), &args); + assert!( + err.contains("skipMemory") && err.contains("--skip-memory"), + "{err}" + ); + assert_no_request_attempted(&err); + let _ = fs::remove_dir_all(&home); +} + +#[test] +fn task_feedback_validates_its_arguments_locally() { + let home = temp_home(); + seed_credentials(&home); + + let args = [ + "agent", + "task", + "feedback", + "agt-1", + "run-1", + "--workspace", + "ws-1", + "--rating", + "meh", + ]; + let err = assert_failure(&run(&home, &args), &args); + assert!(err.contains("up") && err.contains("down"), "{err}"); + + let long = "x".repeat(2001); + let args = [ + "agent", + "task", + "feedback", + "agt-1", + "run-1", + "--workspace", + "ws-1", + "--rating", + "up", + "--comment", + long.as_str(), + ]; + let err = assert_failure(&run(&home, &args), &args); + assert!(err.contains("2001 characters"), "{err}"); + assert_no_request_attempted(&err); + + let _ = fs::remove_dir_all(&home); +} + +#[test] +fn task_list_rejects_a_page_size_over_the_api_limit() { + let home = temp_home(); + let args = [ + "agent", + "task", + "list", + "agt-1", + "--workspace", + "ws-1", + "--page-size", + "101", + ]; + let err = assert_failure(&run(&home, &args), &args); + assert!(err.contains("101"), "{err}"); + let _ = fs::remove_dir_all(&home); +} diff --git a/crates/cli/tests/agent/wire.rs b/crates/cli/tests/agent/wire.rs new file mode 100644 index 0000000..809a609 --- /dev/null +++ b/crates/cli/tests/agent/wire.rs @@ -0,0 +1,426 @@ +//! Wire-level tests for the A2A `agent` subcommands (`card`, `send`, `task`) +//! against a loopback stub. +//! +//! The stub answers bare A2A JSON, not a MemoryLake envelope, exactly as the +//! production endpoints do. These pin the HTTP method, path, query string, +//! body and headers each subcommand sends, and what it prints back. + +use serde_json::Value; + +use crate::common::stub::{ + exchange_event_stream_with_remembered_workspace, exchange_with_remembered_workspace, + request_body, request_header, request_line, +}; +use crate::common::{assert_failure, assert_success}; + +const WS: &str = "ws-1"; +const AGENT: &str = "agt-1"; +const ROOT: &str = "/api/v3/workspaces/ws-1/agents/agt-1"; + +const CARD: &str = + r#"{"protocolVersion":"0.3","name":"Default Agent","capabilities":{"streaming":true}}"#; +const COMPLETED_TASK: &str = r#"{"task":{"id":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_COMPLETED"},"artifacts":[{"artifactId":"result-run-1","name":"result","parts":[{"text":"PONG"}]}],"history":[{"role":"ROLE_USER","parts":[{"text":"ping"}]},{"role":"ROLE_AGENT","parts":[{"text":"PONG"}]}]}}"#; +const WORKING_TASK: &str = r#"{"task":{"id":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_WORKING"},"history":[]}}"#; +const INPUT_REQUIRED_TASK: &str = r#"{"task":{"id":"run-2","contextId":"ctx-1","status":{"state":"TASK_STATE_INPUT_REQUIRED","message":{"role":"ROLE_AGENT","parts":[{"text":"Which file?"}]}},"artifacts":[]}}"#; +const BARE_TASK: &str = r#"{"id":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_COMPLETED"},"metadata":{"task-feedback/v1":{"rating":"up"}}}"#; +const TASK_PAGE: &str = r#"{"tasks":[],"pageSize":5,"totalSize":0}"#; + +fn body_json(request: &str) -> Value { + serde_json::from_str(request_body(request)).expect("request body is JSON") +} + +fn stdout_of(output: &std::process::Output) -> String { + String::from_utf8_lossy(&output.stdout).into_owned() +} + +fn stderr_of(output: &std::process::Output) -> String { + String::from_utf8_lossy(&output.stderr).into_owned() +} + +#[test] +fn card_gets_the_well_known_document_and_prints_it_as_is() { + let args = ["agent", "card", AGENT]; + let (request, output) = exchange_with_remembered_workspace(CARD, WS, &args); + assert_success(&output, &args); + + assert_eq!( + request_line(&request), + format!("GET {ROOT}/.well-known/agent-card.json HTTP/1.1") + ); + let printed: Value = serde_json::from_str(&stdout_of(&output)).expect("card JSON"); + assert_eq!(printed["name"], "Default Agent"); +} + +#[test] +fn send_posts_a_v1_message_and_prints_the_reply_text() { + let args = ["agent", "send", AGENT, "--text", "ping"]; + let (request, output) = exchange_with_remembered_workspace(COMPLETED_TASK, WS, &args); + assert_success(&output, &args); + + assert_eq!( + request_line(&request), + format!("POST {ROOT}/a2a/message:send HTTP/1.1"), + "v1.0 lives on the unversioned path; `/a2a/v1/` is protocol 0.3" + ); + let body = body_json(&request); + assert_eq!(body["message"]["role"], "ROLE_USER"); + assert_eq!( + body["message"]["parts"], + serde_json::json!([{"text": "ping"}]) + ); + assert!( + body["message"]["messageId"] + .as_str() + .is_some_and(|id| !id.is_empty()), + "{body}" + ); + assert!( + body.get("configuration").is_none(), + "default send is blocking: {body}" + ); + assert!( + body.get("metadata").is_none(), + "no scope flags, no metadata: {body}" + ); + + assert_eq!( + stdout_of(&output), + "PONG\n", + "stdout carries only the reply" + ); + let stderr = stderr_of(&output); + assert!(stderr.contains("task run-1"), "{stderr}"); + assert!(stderr.contains("context ctx-1"), "{stderr}"); + assert!(stderr.contains("TASK_STATE_COMPLETED"), "{stderr}"); +} + +#[test] +fn send_maps_scope_flags_onto_the_memorylake_extension() { + let args = [ + "agent", + "send", + AGENT, + "--text", + "ping", + "--context", + "ctx-9", + "--task", + "run-9", + "--actor", + "act-1", + "--project", + "proj-rw", + "--read-only-project", + "proj-a", + "--read-only-project", + "proj-b", + "--skip-memory", + "--metadata-json", + r#"{"overrides":{"maxTurns":2}}"#, + "--raw", + ]; + let (request, output) = exchange_with_remembered_workspace(COMPLETED_TASK, WS, &args); + assert_success(&output, &args); + + let body = body_json(&request); + assert_eq!(body["message"]["contextId"], "ctx-9"); + assert_eq!(body["message"]["taskId"], "run-9"); + assert_eq!( + body["metadata"]["memorylake"], + serde_json::json!({ + "actorId": "act-1", + "readWriteProjectId": "proj-rw", + "readOnlyProjectIds": ["proj-a", "proj-b"], + "skipMemory": true, + "overrides": {"maxTurns": 2} + }) + ); + + let printed: Value = serde_json::from_str(&stdout_of(&output)).expect("--raw prints JSON"); + assert_eq!(printed["task"]["id"], "run-1"); +} + +#[test] +fn send_no_wait_asks_for_an_immediate_return_and_prints_the_task() { + let args = ["agent", "send", AGENT, "--text", "ping", "--no-wait"]; + let (request, output) = exchange_with_remembered_workspace(WORKING_TASK, WS, &args); + assert_success(&output, &args); + + let body = body_json(&request); + assert_eq!( + body["configuration"], + serde_json::json!({"returnImmediately": true}) + ); + + let printed: Value = serde_json::from_str(&stdout_of(&output)).expect("task JSON"); + assert_eq!(printed["task"]["status"]["state"], "TASK_STATE_WORKING"); +} + +#[test] +fn send_explains_how_to_answer_an_input_required_task() { + let args = ["agent", "send", AGENT, "--text", "open it"]; + let (_, output) = exchange_with_remembered_workspace(INPUT_REQUIRED_TASK, WS, &args); + assert_success(&output, &args); + + assert_eq!(stdout_of(&output), "Which file?\n"); + let stderr = stderr_of(&output); + assert!( + stderr.contains("--task run-2") && stderr.contains("--context ctx-1"), + "{stderr}" + ); +} + +#[test] +fn send_accepts_full_parts_from_json() { + let args = [ + "agent", + "send", + AGENT, + "--message-json", + r#"[{"text":"look"},{"url":"https://example.com/x.png","mediaType":"image/png"}]"#, + ]; + let (request, output) = exchange_with_remembered_workspace(COMPLETED_TASK, WS, &args); + assert_success(&output, &args); + + let body = body_json(&request); + assert_eq!(body["message"]["parts"][1]["mediaType"], "image/png"); +} + +#[test] +fn send_stream_prints_fragments_as_they_arrive_and_the_artifact_only_as_fallback() { + // The production sequence: task, then token-by-token status updates, then + // the whole reply again as an artifact, then the terminal status. + let events = [ + r#"{"task":{"id":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_WORKING"}}}"#, + r#"{"statusUpdate":{"taskId":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_WORKING"}}}"#, + r#"{"statusUpdate":{"taskId":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_WORKING","message":{"parts":[{"text":"PO"}]}}}}"#, + r#"{"statusUpdate":{"taskId":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_WORKING","message":{"parts":[{"text":"NG"}]}}}}"#, + r#"{"artifactUpdate":{"taskId":"run-1","contextId":"ctx-1","lastChunk":true,"artifact":{"parts":[{"text":"PONG"}]}}}"#, + r#"{"statusUpdate":{"taskId":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_COMPLETED"}}}"#, + ]; + let args = ["agent", "send", AGENT, "--text", "ping", "--stream"]; + let (request, output) = exchange_event_stream_with_remembered_workspace(&events, WS, &args); + assert_success(&output, &args); + + assert_eq!( + request_line(&request), + format!("POST {ROOT}/a2a/message:stream HTTP/1.1") + ); + assert_eq!( + request_header(&request, "accept").map(str::to_ascii_lowercase), + Some("text/event-stream".into()), + "{request}" + ); + + assert_eq!( + stdout_of(&output), + "PONG\n", + "the artifact repeats the streamed text and must not be printed twice" + ); + let stderr = stderr_of(&output); + assert!( + stderr.contains("task run-1") && stderr.contains("TASK_STATE_COMPLETED"), + "{stderr}" + ); +} + +#[test] +fn send_stream_falls_back_to_the_artifact_when_nothing_was_streamed() { + let events = [ + r#"{"task":{"id":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_WORKING"}}}"#, + r#"{"artifactUpdate":{"taskId":"run-1","contextId":"ctx-1","artifact":{"parts":[{"text":"PONG"}]}}}"#, + r#"{"statusUpdate":{"taskId":"run-1","contextId":"ctx-1","status":{"state":"TASK_STATE_COMPLETED"}}}"#, + ]; + let args = ["agent", "send", AGENT, "--text", "ping", "--stream"]; + let (_, output) = exchange_event_stream_with_remembered_workspace(&events, WS, &args); + assert_success(&output, &args); + assert_eq!(stdout_of(&output), "PONG\n"); +} + +#[test] +fn send_stream_raw_prints_one_event_per_line() { + let events = [ + r#"{"task":{"id":"run-1"}}"#, + r#"{"statusUpdate":{"taskId":"run-1","status":{"state":"TASK_STATE_COMPLETED"}}}"#, + ]; + let args = [ + "agent", "send", AGENT, "--text", "ping", "--stream", "--raw", + ]; + let (_, output) = exchange_event_stream_with_remembered_workspace(&events, WS, &args); + assert_success(&output, &args); + + let lines: Vec = stdout_of(&output) + .lines() + .map(|line| serde_json::from_str(line).expect("each line is one JSON event")) + .collect(); + assert_eq!(lines.len(), 2); + assert_eq!(lines[0]["task"]["id"], "run-1"); + assert!(stderr_of(&output).is_empty(), "raw mode adds no summary"); +} + +#[test] +fn task_list_sends_camel_case_query_parameters() { + let args = [ + "agent", + "task", + "list", + AGENT, + "--context", + "ctx-1", + "--status", + "TASK_STATE_WORKING", + "--page-size", + "5", + "--page-token", + "tok", + "--after", + "2026-09-01T00:00:00Z", + "--history-length", + "2", + "--artifacts", + ]; + let (request, output) = exchange_with_remembered_workspace(TASK_PAGE, WS, &args); + assert_success(&output, &args); + + let line = request_line(&request); + assert!( + line.starts_with(&format!("GET {ROOT}/a2a/tasks?")), + "{line}" + ); + for expected in [ + "contextId=ctx-1", + "status=TASK_STATE_WORKING", + "pageSize=5", + "pageToken=tok", + "statusTimestampAfter=2026-09-01T00%3A00%3A00Z", + "historyLength=2", + "includeArtifacts=true", + ] { + assert!(line.contains(expected), "missing {expected} in {line}"); + } +} + +#[test] +fn task_list_without_filters_sends_no_query() { + let args = ["agent", "task", "list", AGENT]; + let (request, output) = exchange_with_remembered_workspace(TASK_PAGE, WS, &args); + assert_success(&output, &args); + assert_eq!( + request_line(&request), + format!("GET {ROOT}/a2a/tasks HTTP/1.1") + ); +} + +#[test] +fn task_get_activates_the_feedback_extension() { + let args = [ + "agent", + "task", + "get", + AGENT, + "run-1", + "--history-length", + "3", + ]; + let (request, output) = exchange_with_remembered_workspace(BARE_TASK, WS, &args); + assert_success(&output, &args); + + assert_eq!( + request_line(&request), + format!("GET {ROOT}/a2a/tasks/run-1?historyLength=3 HTTP/1.1") + ); + assert_eq!( + request_header(&request, "a2a-extensions"), + Some("extensions://task-feedback/v1"), + "without the header the server hides the rating: {request}" + ); + let printed: Value = serde_json::from_str(&stdout_of(&output)).expect("task JSON"); + assert_eq!(printed["metadata"]["task-feedback/v1"]["rating"], "up"); +} + +#[test] +fn task_cancel_posts_to_the_cancel_verb() { + let args = ["agent", "task", "cancel", AGENT, "run-1"]; + let (request, output) = exchange_with_remembered_workspace(BARE_TASK, WS, &args); + assert_success(&output, &args); + assert_eq!( + request_line(&request), + format!("POST {ROOT}/a2a/tasks/run-1:cancel HTTP/1.1") + ); +} + +#[test] +fn task_feedback_posts_the_rating_with_the_extension_header() { + let args = [ + "agent", + "task", + "feedback", + AGENT, + "run-1", + "--rating", + "down", + "--comment", + "meh", + ]; + let (request, output) = exchange_with_remembered_workspace(BARE_TASK, WS, &args); + assert_success(&output, &args); + + assert_eq!( + request_line(&request), + format!("POST {ROOT}/a2a/tasks/run-1:feedback HTTP/1.1") + ); + assert_eq!( + request_header(&request, "a2a-extensions"), + Some("extensions://task-feedback/v1"), + "{request}" + ); + assert_eq!( + body_json(&request), + serde_json::json!({"rating": "down", "comment": "meh"}) + ); +} + +#[test] +fn task_feedback_without_a_comment_omits_the_field() { + let args = [ + "agent", "task", "feedback", AGENT, "run-1", "--rating", "up", + ]; + let (request, output) = exchange_with_remembered_workspace(BARE_TASK, WS, &args); + assert_success(&output, &args); + assert_eq!(body_json(&request), serde_json::json!({"rating": "up"})); +} + +#[test] +fn an_explicit_workspace_flag_overrides_the_remembered_one() { + let args = ["agent", "card", AGENT, "--workspace", "ws-other"]; + let (request, output) = exchange_with_remembered_workspace(CARD, WS, &args); + assert_success(&output, &args); + assert!( + request_line(&request).contains("/workspaces/ws-other/"), + "{}", + request_line(&request) + ); +} + +#[test] +fn an_a2a_error_is_reported_by_its_reason_not_as_escaped_json() { + // Production's answer to cancelling a finished task: a MemoryLake + // envelope carrying the A2A error document as a string. The stub answers + // 200, so this exercises the `success: false` path of the decoder. + let body = r#"{"success":false,"message":"{\"error\":{\"code\":400,\"status\":\"FAILED_PRECONDITION\",\"message\":\"Task cannot be canceled - current state: 3\",\"details\":[{\"reason\":\"TASK_NOT_CANCELABLE\",\"domain\":\"a2a-protocol.org\"}]}}","error_code":"INVALID_ARGUMENT"}"#; + let args = ["agent", "task", "cancel", AGENT, "run-1"]; + let (_, output) = exchange_with_remembered_workspace(body, WS, &args); + let err = assert_failure(&output, &args); + assert!( + err.contains("Task cannot be canceled - current state: 3 (TASK_NOT_CANCELABLE)"), + "{err}" + ); + let headline = err + .lines() + .find(|line| line.contains("Task cannot be canceled")) + .expect("headline names the failure"); + assert!( + !headline.contains("\\\"error\\\""), + "the escaped document must not be the headline: {err}" + ); +} diff --git a/crates/cli/tests/common/stub.rs b/crates/cli/tests/common/stub.rs index cb0871f..46508fd 100644 --- a/crates/cli/tests/common/stub.rs +++ b/crates/cli/tests/common/stub.rs @@ -69,6 +69,46 @@ impl StubServer { } } + /// Answer one request with a `text/event-stream`, writing `events` one + /// frame at a time with a pause in between. + /// + /// The pauses matter: they make each frame arrive in its own read, so a + /// command that only works when the whole body is buffered fails here. + fn event_stream(events: &[&str]) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind stub server"); + let addr = listener.local_addr().expect("stub server address"); + let frames: Vec = events + .iter() + .map(|event| format!("data:{event}\n\n")) + .collect(); + + let (sender, requests) = channel(); + let handle = std::thread::spawn(move || { + let Ok((mut stream, _)) = listener.accept() else { + return; + }; + let request = read_http_request(&mut stream); + let _ = sender.send(request); + let head = + "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\nconnection: close\r\n\r\n"; + let _ = stream.write_all(head.as_bytes()); + let _ = stream.flush(); + for frame in frames { + std::thread::sleep(Duration::from_millis(20)); + let _ = stream.write_all(frame.as_bytes()); + let _ = stream.flush(); + } + // No content-length: the body ends when the connection does. + let _ = stream.shutdown(std::net::Shutdown::Write); + }); + + Self { + base_url: format!("http://{addr}"), + requests, + handle: Some(handle), + } + } + fn received(&self) -> String { match self.requests.recv_timeout(Duration::from_secs(10)) { Ok(request) => request, @@ -232,3 +272,33 @@ pub fn request_body(request: &str) -> &str { .map(|(_, body)| body) .unwrap_or_default() } + +/// Run one command against a stub that answers with a server-sent event +/// stream of `events` (one JSON document each), from a `$HOME` that already +/// remembers `workspace`. +pub fn exchange_event_stream_with_remembered_workspace( + events: &[&str], + workspace: &str, + args: &[&str], +) -> (String, Output) { + let server = StubServer::event_stream(events); + let home = logged_in_home_with_workspace(&server.base_url, workspace); + let output = run(&home, args); + let request = server.received(); + let _ = fs::remove_dir_all(&home); + (request, output) +} + +/// A header's value from a raw HTTP request, matched case-insensitively. +pub fn request_header<'a>(request: &'a str, name: &str) -> Option<&'a str> { + let needle = format!("{}:", name.to_ascii_lowercase()); + request + .lines() + .skip(1) + .take_while(|line| !line.is_empty()) + .find_map(|line| { + line.to_ascii_lowercase() + .starts_with(&needle) + .then(|| line[needle.len()..].trim()) + }) +} diff --git a/crates/core/src/api/agents/a2a/card.rs b/crates/core/src/api/agents/a2a/card.rs new file mode 100644 index 0000000..9c0c015 --- /dev/null +++ b/crates/core/src/api/agents/a2a/card.rs @@ -0,0 +1,49 @@ +//! Fetch an agent's A2A card +//! (`GET .../agents/{agent}/.well-known/agent-card.json`). + +use serde_json::Value; + +use crate::client::Client; +use crate::error::Result; + +use super::{agent_card_path, describe_a2a_error}; + +/// Fetch the agent card: name, supported protocol bindings, capabilities and +/// the extensions the agent understands. +/// +/// Returned as the server sends it. The card's `protocolVersion` reports the +/// default binding's version and is not the version this module speaks; the +/// per-binding truth is in `supportedInterfaces`. +pub fn get_agent_card(client: &Client, workspace_id: &str, agent_id: &str) -> Result { + client + .get_json_with_headers(&agent_card_path(workspace_id, agent_id), &[], &[]) + .map_err(describe_a2a_error) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::{json_ok, one_shot_server}; + + #[test] + fn returns_the_card_verbatim_without_looking_for_an_envelope() { + // The card is bare JSON with no `success` key. + let card = + r#"{"protocolVersion":"0.3","name":"Default Agent","capabilities":{"streaming":true}}"#; + let (base_url, server) = one_shot_server(json_ok(card)); + let client = Client::new(base_url, "sk-test").unwrap(); + + let value = get_agent_card(&client, "ws-1", "agt-1").expect("card"); + assert_eq!(value["name"], "Default Agent"); + assert_eq!(value["capabilities"]["streaming"], true); + + let captured = server.join().unwrap(); + assert!( + captured.head.starts_with( + "GET /api/v3/workspaces/ws-1/agents/agt-1/.well-known/agent-card.json " + ), + "{}", + captured.head + ); + } +} diff --git a/crates/core/src/api/agents/a2a/error.rs b/crates/core/src/api/agents/a2a/error.rs new file mode 100644 index 0000000..cce4106 --- /dev/null +++ b/crates/core/src/api/agents/a2a/error.rs @@ -0,0 +1,157 @@ +//! Making A2A errors readable. +//! +//! An A2A failure reaches the client as a MemoryLake envelope whose `message` +//! is the A2A error document serialized *as a string*: +//! +//! ```json +//! {"success":false,"error_code":"NOT_FOUND", +//! "message":"{\"error\":{\"code\":404,\"status\":\"NOT_FOUND\",\"message\":\"Task not found\", +//! \"details\":[{\"reason\":\"TASK_NOT_FOUND\",\"domain\":\"a2a-protocol.org\"}]}}"} +//! ``` +//! +//! The generic envelope decoder puts that string, escapes and all, at the top +//! of the error. This pass replaces it with the inner message and reason. + +use serde::Deserialize; + +use crate::error::Error; + +#[derive(Debug, Deserialize)] +struct A2aErrorDocument { + error: A2aErrorBody, +} + +#[derive(Debug, Deserialize)] +struct A2aErrorBody { + #[serde(default)] + message: Option, + #[serde(default)] + details: Vec, +} + +#[derive(Debug, Deserialize)] +struct A2aErrorDetail { + #[serde(default)] + reason: Option, +} + +/// Rewrite the first line of an API error when it is an A2A error document. +/// +/// `Task cannot be canceled - current state: 3 (TASK_NOT_CANCELABLE)` replaces +/// the escaped JSON; the envelope's `error_code` and the HTTP transcript that +/// follow are kept. Errors of any other shape pass through untouched. +pub fn describe_a2a_error(err: Error) -> Error { + let Error::Api { message, code } = err else { + return err; + }; + let Some((first, rest)) = split_first_line(&message) else { + return Error::Api { message, code }; + }; + let Some(summary) = summarize(first) else { + return Error::Api { message, code }; + }; + let message = match rest { + Some(rest) => format!("{summary}\n{rest}"), + None => summary, + }; + Error::Api { message, code } +} + +fn split_first_line(message: &str) -> Option<(&str, Option<&str>)> { + match message.split_once('\n') { + Some((first, rest)) => Some((first, Some(rest))), + None if !message.is_empty() => Some((message, None)), + None => None, + } +} + +/// ` ()` if `line` is an A2A error document, possibly +/// followed by the ` [ERROR_CODE]` suffix the envelope decoder appends. +fn summarize(line: &str) -> Option { + let (json, suffix) = match line.rsplit_once(" [") { + Some((json, tail)) if tail.ends_with(']') && json.ends_with('}') => { + (json, Some(&line[json.len()..])) + } + _ => (line, None), + }; + let document: A2aErrorDocument = serde_json::from_str(json).ok()?; + + let message = document + .error + .message + .filter(|m| !m.trim().is_empty()) + .unwrap_or_else(|| "A2A request failed".to_string()); + let reason = document + .error + .details + .iter() + .find_map(|detail| detail.reason.as_deref().filter(|r| !r.trim().is_empty())); + + let mut summary = message; + if let Some(reason) = reason { + summary.push_str(&format!(" ({reason})")); + } + if let Some(suffix) = suffix { + summary.push_str(suffix); + } + Some(summary) +} + +#[cfg(test)] +mod tests { + use super::*; + + const DOC: &str = r#"{"error":{"code":404,"status":"NOT_FOUND","message":"Task not found","details":[{"@type":"type.googleapis.com/google.rpc.ErrorInfo","reason":"TASK_NOT_FOUND","domain":"a2a-protocol.org","metadata":{}}]}}"#; + + fn api(message: String) -> Error { + Error::Api { + message, + code: Some("NOT_FOUND".into()), + } + } + + #[test] + fn replaces_the_escaped_document_with_message_and_reason() { + let err = describe_a2a_error(api(format!("{DOC} [NOT_FOUND]\nHTTP 404\n{{...}}"))); + assert_eq!( + err.to_string(), + "Task not found (TASK_NOT_FOUND) [NOT_FOUND]\nHTTP 404\n{...}" + ); + assert!(matches!(err, Error::Api { code: Some(c), .. } if c == "NOT_FOUND")); + } + + #[test] + fn works_without_the_code_suffix_or_a_transcript() { + assert_eq!( + describe_a2a_error(api(DOC.to_string())).to_string(), + "Task not found (TASK_NOT_FOUND)" + ); + } + + #[test] + fn a_document_without_details_keeps_just_the_message() { + let doc = r#"{"error":{"code":500,"message":"boom"}}"#; + assert_eq!(describe_a2a_error(api(doc.to_string())).to_string(), "boom"); + } + + #[test] + fn ordinary_errors_pass_through_unchanged() { + let plain = "You don't have the \"Chat (A2A)\" permission [ACCESS_DENIED]\nHTTP 403"; + assert_eq!( + describe_a2a_error(api(plain.to_string())).to_string(), + plain + ); + + let not_api = Error::NotLoggedIn; + assert!(matches!(describe_a2a_error(not_api), Error::NotLoggedIn)); + } + + #[test] + fn json_that_is_not_an_a2a_document_passes_through() { + let other = r#"{"foo":"bar"} [X]"#; + assert_eq!( + describe_a2a_error(api(other.to_string())).to_string(), + other + ); + } +} diff --git a/crates/core/src/api/agents/a2a/mod.rs b/crates/core/src/api/agents/a2a/mod.rs new file mode 100644 index 0000000..79803f3 --- /dev/null +++ b/crates/core/src/api/agents/a2a/mod.rs @@ -0,0 +1,153 @@ +//! Talking to a bound agent over A2A +//! (`/api/v3/workspaces/{workspace}/agents/{agent}/a2a/...`). +//! +//! MemoryLake exposes every agent bound to a workspace as an A2A server. This +//! module binds the **A2A v1.0 HTTP+JSON** surface only: the JSON-RPC endpoint +//! is the same operations behind one URL, and the v0.3 REST surface is a +//! superseded protocol revision. +//! +//! Mind the paths: the v1.0 REST operations live on *unversioned* paths +//! (`/a2a/message:send`, `/a2a/tasks`), while the paths carrying `/v1/` +//! (`/a2a/v1/message:send`) belong to protocol version 0.3. The agent card's +//! `supportedInterfaces` says so explicitly; the constants here are named by +//! protocol version, not by path. +//! +//! Responses are A2A-shaped, not MemoryLake envelopes, and the API documents +//! them without a schema, so they come back as [`serde_json::Value`]. Errors +//! do arrive as envelopes and are decoded like every other API error, then +//! run through [`describe_a2a_error`] to surface the A2A reason. + +mod card; +mod error; +mod send; +mod tasks; +mod types; + +pub use card::get_agent_card; +pub use error::describe_a2a_error; +pub use send::{send_message, stream_message}; +pub use tasks::{ + FEEDBACK_COMMENT_MAX_CHARS, ListTasksParams, Rating, TaskFeedbackRequest, cancel_task, + get_task, list_tasks, submit_task_feedback, +}; +pub use types::{ + MemorylakeExtension, Message, ROLE_USER, SendConfiguration, SendMessageRequest, SendMetadata, + TASK_STATE_CANCELED, TASK_STATE_COMPLETED, TASK_STATE_FAILED, TASK_STATE_INPUT_REQUIRED, + TASK_STATE_REJECTED, TASK_STATE_WORKING, is_terminal_state, text_part, +}; + +use crate::api::path::encode_segment; + +use super::workspace_agent_path; + +/// Header that activates A2A extensions for a request. +/// +/// Required to read task feedback back: without it `GET tasks/{id}` omits the +/// `metadata` the rating lives in. Submitting feedback works either way, so +/// the header is sent on both to keep the two symmetric. +pub const A2A_EXTENSIONS_HEADER: &str = "A2A-Extensions"; + +/// URI of MemoryLake's task-feedback extension, as advertised in the agent +/// card's `capabilities.extensions`. +pub const TASK_FEEDBACK_EXTENSION_URI: &str = "extensions://task-feedback/v1"; + +/// `/api/v3/workspaces/{ws}/agents/{agent}/a2a` +fn a2a_root(workspace_id: &str, agent_id: &str) -> String { + format!("{}/a2a", workspace_agent_path(workspace_id, agent_id)) +} + +/// `/api/v3/workspaces/{ws}/agents/{agent}/.well-known/agent-card.json` +fn agent_card_path(workspace_id: &str, agent_id: &str) -> String { + format!( + "{}/.well-known/agent-card.json", + workspace_agent_path(workspace_id, agent_id) + ) +} + +/// `.../a2a/message:send` (A2A v1.0) +fn message_send_path(workspace_id: &str, agent_id: &str) -> String { + format!("{}/message:send", a2a_root(workspace_id, agent_id)) +} + +/// `.../a2a/message:stream` (A2A v1.0) +fn message_stream_path(workspace_id: &str, agent_id: &str) -> String { + format!("{}/message:stream", a2a_root(workspace_id, agent_id)) +} + +/// `.../a2a/tasks` (A2A v1.0) +fn tasks_path(workspace_id: &str, agent_id: &str) -> String { + format!("{}/tasks", a2a_root(workspace_id, agent_id)) +} + +/// `.../a2a/tasks/{task}` (A2A v1.0) +fn task_path(workspace_id: &str, agent_id: &str, task_id: &str) -> String { + format!( + "{}/{}", + tasks_path(workspace_id, agent_id), + encode_segment(task_id) + ) +} + +/// `.../a2a/tasks/{task}:cancel` (A2A v1.0) +fn task_cancel_path(workspace_id: &str, agent_id: &str, task_id: &str) -> String { + format!("{}:cancel", task_path(workspace_id, agent_id, task_id)) +} + +/// `.../a2a/tasks/{task}:feedback` (MemoryLake extension, A2A v1.0 only) +fn task_feedback_path(workspace_id: &str, agent_id: &str, task_id: &str) -> String { + format!("{}:feedback", task_path(workspace_id, agent_id, task_id)) +} + +/// Headers that turn the feedback extension on for a request. +fn feedback_extension_headers() -> [(&'static str, &'static str); 1] { + [(A2A_EXTENSIONS_HEADER, TASK_FEEDBACK_EXTENSION_URI)] +} + +#[cfg(test)] +mod tests { + use super::*; + + const ROOT: &str = "/api/v3/workspaces/ws-1/agents/agt-1"; + + #[test] + fn v1_paths_carry_no_version_segment() { + // The `/v1/` paths are protocol 0.3; v1.0 is the unversioned set. + assert_eq!( + message_send_path("ws-1", "agt-1"), + format!("{ROOT}/a2a/message:send") + ); + assert_eq!( + message_stream_path("ws-1", "agt-1"), + format!("{ROOT}/a2a/message:stream") + ); + assert_eq!(tasks_path("ws-1", "agt-1"), format!("{ROOT}/a2a/tasks")); + assert_eq!( + task_path("ws-1", "agt-1", "run-1"), + format!("{ROOT}/a2a/tasks/run-1") + ); + assert_eq!( + task_cancel_path("ws-1", "agt-1", "run-1"), + format!("{ROOT}/a2a/tasks/run-1:cancel") + ); + assert_eq!( + task_feedback_path("ws-1", "agt-1", "run-1"), + format!("{ROOT}/a2a/tasks/run-1:feedback") + ); + } + + #[test] + fn agent_card_lives_under_well_known() { + assert_eq!( + agent_card_path("ws-1", "agt-1"), + format!("{ROOT}/.well-known/agent-card.json") + ); + } + + #[test] + fn a_task_id_cannot_escape_its_segment() { + assert_eq!( + task_path("ws-1", "agt-1", "run-1/../x?y"), + format!("{ROOT}/a2a/tasks/run-1%2F..%2Fx%3Fy") + ); + } +} diff --git a/crates/core/src/api/agents/a2a/send.rs b/crates/core/src/api/agents/a2a/send.rs new file mode 100644 index 0000000..d821800 --- /dev/null +++ b/crates/core/src/api/agents/a2a/send.rs @@ -0,0 +1,149 @@ +//! Send a message to an agent (`POST .../a2a/message:send` and +//! `message:stream`, A2A v1.0). + +use std::io::BufReader; + +use serde_json::Value; + +use crate::client::Client; +use crate::error::Result; +use crate::sse::EventStream; + +use super::types::SendMessageRequest; +use super::{describe_a2a_error, message_send_path, message_stream_path}; + +/// Deliver a message and return the resulting task. +/// +/// The server blocks until the task finishes unless +/// `configuration.returnImmediately` is set. The response is +/// `{"task": {...}}` — or, for a reply that produced no task, `{"message": +/// {...}}` — exactly as the protocol defines it. +pub fn send_message( + client: &Client, + workspace_id: &str, + agent_id: &str, + request: &SendMessageRequest, +) -> Result { + client + .post_json_with_headers(&message_send_path(workspace_id, agent_id), request, &[]) + .map_err(describe_a2a_error) +} + +/// Deliver a message and follow the task as a stream of events. +/// +/// Each event is one A2A `StreamResponse`: the first carries `task`, the +/// following ones `statusUpdate` (whose `status.message.parts` hold the +/// agent's text as it is produced) or `artifactUpdate`, and the last has a +/// terminal `status.state`. +pub fn stream_message( + client: &Client, + workspace_id: &str, + agent_id: &str, + request: &SendMessageRequest, +) -> Result>> { + client + .post_event_stream(&message_stream_path(workspace_id, agent_id), request, &[]) + .map_err(describe_a2a_error) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::api::agents::a2a::{Message, ROLE_USER, text_part}; + use crate::test_support::{json_ok, one_shot_server}; + + fn request() -> SendMessageRequest { + SendMessageRequest { + message: Message { + role: ROLE_USER.into(), + message_id: "m-1".into(), + parts: vec![text_part("hi")], + context_id: None, + task_id: None, + }, + configuration: None, + metadata: None, + } + } + + #[test] + fn send_posts_to_the_unversioned_v1_path() { + let task = r#"{"task":{"id":"run-1","status":{"state":"TASK_STATE_COMPLETED"}}}"#; + let (base_url, server) = one_shot_server(json_ok(task)); + let client = Client::new(base_url, "sk-test").unwrap(); + + let value = send_message(&client, "ws-1", "agt-1", &request()).expect("send"); + assert_eq!(value["task"]["id"], "run-1"); + + let captured = server.join().unwrap(); + assert!( + captured + .head + .starts_with("POST /api/v3/workspaces/ws-1/agents/agt-1/a2a/message:send "), + "{}", + captured.head + ); + let body: Value = serde_json::from_slice(&captured.body).unwrap(); + assert_eq!(body["message"]["parts"][0]["text"], "hi"); + } + + #[test] + fn stream_reads_events_until_the_body_ends() { + let body = "data:{\"task\":{\"id\":\"run-1\"}}\n\ndata:{\"statusUpdate\":{\"status\":{\"state\":\"TASK_STATE_COMPLETED\"}}}\n\n"; + let response = format!( + "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\n\r\n{body}", + body.len() + ); + let (base_url, server) = one_shot_server(response); + let client = Client::new(base_url, "sk-test").unwrap(); + + let events: Vec = stream_message(&client, "ws-1", "agt-1", &request()) + .expect("open stream") + .collect::>() + .expect("read events"); + assert_eq!(events.len(), 2); + assert_eq!(events[0]["task"]["id"], "run-1"); + assert_eq!( + events[1]["statusUpdate"]["status"]["state"], + "TASK_STATE_COMPLETED" + ); + + let captured = server.join().unwrap(); + assert!( + captured + .head + .starts_with("POST /api/v3/workspaces/ws-1/agents/agt-1/a2a/message:stream "), + "{}", + captured.head + ); + assert!(captured.has_header("accept"), "{}", captured.head); + } + + #[test] + fn stream_answered_with_a_json_envelope_is_the_envelope_error() { + // Observed in production: a permission failure on `message:stream` + // comes back as a plain envelope, not as an event stream. + let body = r#"{"success":false,"message":"You don't have the \"Chat (A2A)\" permission","error_code":"ACCESS_DENIED"}"#; + let (base_url, _server) = one_shot_server(format!( + "HTTP/1.1 403 Forbidden\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{body}", + body.len() + )); + let client = Client::new(base_url, "sk-test").unwrap(); + + let err = stream_message(&client, "ws-1", "agt-1", &request()).expect_err("denied"); + assert!(err.to_string().contains("Chat (A2A)"), "{err}"); + assert!( + matches!(&err, crate::Error::Api { code: Some(code), .. } if code == "ACCESS_DENIED"), + "{err:?}" + ); + } + + #[test] + fn a_2xx_that_is_not_an_event_stream_is_rejected() { + let (base_url, _server) = one_shot_server(json_ok(r#"{"success":true,"data":{}}"#)); + let client = Client::new(base_url, "sk-test").unwrap(); + + let err = stream_message(&client, "ws-1", "agt-1", &request()).expect_err("not a stream"); + assert!(err.to_string().contains("text/event-stream"), "{err}"); + } +} diff --git a/crates/core/src/api/agents/a2a/tasks.rs b/crates/core/src/api/agents/a2a/tasks.rs new file mode 100644 index 0000000..ebd355b --- /dev/null +++ b/crates/core/src/api/agents/a2a/tasks.rs @@ -0,0 +1,326 @@ +//! A2A task operations (`.../a2a/tasks`, A2A v1.0). + +use serde::Serialize; +use serde_json::Value; + +use crate::client::Client; +use crate::error::{Error, Result}; + +use super::{ + describe_a2a_error, feedback_extension_headers, task_cancel_path, task_feedback_path, + task_path, tasks_path, +}; + +/// Longest `comment` the feedback extension accepts, in characters. +pub const FEEDBACK_COMMENT_MAX_CHARS: usize = 2000; + +/// Query parameters of `GET .../a2a/tasks`. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct ListTasksParams { + /// Only tasks in this conversation thread. + pub context_id: Option, + /// Only tasks in this state (`TASK_STATE_...`). + pub status: Option, + /// Tasks per page, 1–100. + pub page_size: Option, + /// `nextPageToken` from the previous page. + pub page_token: Option, + /// How many history messages to include per task. + pub history_length: Option, + /// Only tasks whose status changed after this ISO 8601 instant. + pub status_timestamp_after: Option, + /// Include each task's artifacts. + pub include_artifacts: Option, +} + +impl ListTasksParams { + /// Query pairs, camelCase as A2A spells them. + pub fn to_query(&self) -> Vec<(&'static str, String)> { + let mut query = Vec::new(); + if let Some(value) = &self.context_id { + query.push(("contextId", value.clone())); + } + if let Some(value) = &self.status { + query.push(("status", value.clone())); + } + if let Some(value) = self.page_size { + query.push(("pageSize", value.to_string())); + } + if let Some(value) = &self.page_token { + query.push(("pageToken", value.clone())); + } + if let Some(value) = self.history_length { + query.push(("historyLength", value.to_string())); + } + if let Some(value) = &self.status_timestamp_after { + query.push(("statusTimestampAfter", value.clone())); + } + if let Some(value) = self.include_artifacts { + query.push(("includeArtifacts", value.to_string())); + } + query + } +} + +/// A thumbs-up or thumbs-down on a task. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "lowercase")] +pub enum Rating { + /// The task's result was good. + Up, + /// The task's result was bad. + Down, +} + +impl Rating { + /// Wire spelling. + pub fn as_str(self) -> &'static str { + match self { + Self::Up => "up", + Self::Down => "down", + } + } +} + +/// Body of `POST .../a2a/tasks/{task}:feedback`. +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +pub struct TaskFeedbackRequest { + /// The rating. + pub rating: Rating, + /// Optional free text, at most [`FEEDBACK_COMMENT_MAX_CHARS`] characters. + #[serde(skip_serializing_if = "Option::is_none")] + pub comment: Option, +} + +/// List an agent's tasks, newest first. +/// +/// The response is `{"tasks": [...], "nextPageToken": "...", "pageSize": N, +/// "totalSize": N}`. +pub fn list_tasks( + client: &Client, + workspace_id: &str, + agent_id: &str, + params: &ListTasksParams, +) -> Result { + client + .get_json_with_headers( + &tasks_path(workspace_id, agent_id), + ¶ms.to_query(), + &feedback_extension_headers(), + ) + .map_err(describe_a2a_error) +} + +/// Fetch one task, with `history_length` history messages. +/// +/// The feedback extension is activated so a rating left on the task shows up +/// in its `metadata`. +pub fn get_task( + client: &Client, + workspace_id: &str, + agent_id: &str, + task_id: &str, + history_length: Option, +) -> Result { + let query: Vec<(&str, String)> = history_length + .map(|n| vec![("historyLength", n.to_string())]) + .unwrap_or_default(); + client + .get_json_with_headers( + &task_path(workspace_id, agent_id, task_id), + &query, + &feedback_extension_headers(), + ) + .map_err(describe_a2a_error) +} + +/// Cancel a task that is still running. +/// +/// A task already in a terminal state is refused with reason +/// `TASK_NOT_CANCELABLE`. +pub fn cancel_task( + client: &Client, + workspace_id: &str, + agent_id: &str, + task_id: &str, +) -> Result { + client + .post_json_with_headers( + &task_cancel_path(workspace_id, agent_id, task_id), + &serde_json::json!({}), + &[], + ) + .map_err(describe_a2a_error) +} + +/// Rate a task. Returns the task with the rating in its `metadata`. +pub fn submit_task_feedback( + client: &Client, + workspace_id: &str, + agent_id: &str, + task_id: &str, + request: &TaskFeedbackRequest, +) -> Result { + if let Some(comment) = &request.comment { + let chars = comment.chars().count(); + if chars > FEEDBACK_COMMENT_MAX_CHARS { + return Err(Error::Api { + message: format!( + "feedback comment is {chars} characters; the limit is {FEEDBACK_COMMENT_MAX_CHARS}" + ), + code: None, + }); + } + } + client + .post_json_with_headers( + &task_feedback_path(workspace_id, agent_id, task_id), + request, + &feedback_extension_headers(), + ) + .map_err(describe_a2a_error) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::api::agents::a2a::{A2A_EXTENSIONS_HEADER, TASK_FEEDBACK_EXTENSION_URI}; + use crate::test_support::{json_ok, one_shot_server}; + + fn header_value(head: &str, name: &str) -> Option { + let needle = format!("{}:", name.to_ascii_lowercase()); + head.lines().find_map(|line| { + line.to_ascii_lowercase() + .starts_with(&needle) + .then(|| line[needle.len()..].trim().to_string()) + }) + } + + #[test] + fn list_params_are_camel_case_and_only_set_ones_are_sent() { + assert!(ListTasksParams::default().to_query().is_empty()); + let query = ListTasksParams { + context_id: Some("ctx".into()), + status: Some("TASK_STATE_WORKING".into()), + page_size: Some(5), + page_token: Some("tok".into()), + history_length: Some(2), + status_timestamp_after: Some("2026-09-01T00:00:00Z".into()), + include_artifacts: Some(true), + } + .to_query(); + assert_eq!( + query, + vec![ + ("contextId", "ctx".to_string()), + ("status", "TASK_STATE_WORKING".to_string()), + ("pageSize", "5".to_string()), + ("pageToken", "tok".to_string()), + ("historyLength", "2".to_string()), + ("statusTimestampAfter", "2026-09-01T00:00:00Z".to_string()), + ("includeArtifacts", "true".to_string()), + ] + ); + } + + #[test] + fn get_task_activates_the_feedback_extension() { + let task = r#"{"id":"run-1","status":{"state":"TASK_STATE_COMPLETED"},"metadata":{"task-feedback/v1":{"rating":"up"}}}"#; + let (base_url, server) = one_shot_server(json_ok(task)); + let client = Client::new(base_url, "sk-test").unwrap(); + + let value = get_task(&client, "ws-1", "agt-1", "run-1", Some(3)).expect("task"); + assert_eq!(value["metadata"]["task-feedback/v1"]["rating"], "up"); + + let captured = server.join().unwrap(); + assert!( + captured.head.starts_with( + "GET /api/v3/workspaces/ws-1/agents/agt-1/a2a/tasks/run-1?historyLength=3 " + ), + "{}", + captured.head + ); + assert_eq!( + header_value(&captured.head, A2A_EXTENSIONS_HEADER).as_deref(), + Some(TASK_FEEDBACK_EXTENSION_URI) + ); + } + + #[test] + fn feedback_posts_the_rating_with_the_extension_header() { + let (base_url, server) = one_shot_server(json_ok(r#"{"id":"run-1"}"#)); + let client = Client::new(base_url, "sk-test").unwrap(); + + submit_task_feedback( + &client, + "ws-1", + "agt-1", + "run-1", + &TaskFeedbackRequest { + rating: Rating::Down, + comment: Some("meh".into()), + }, + ) + .expect("feedback"); + + let captured = server.join().unwrap(); + assert!( + captured + .head + .starts_with("POST /api/v3/workspaces/ws-1/agents/agt-1/a2a/tasks/run-1:feedback "), + "{}", + captured.head + ); + assert_eq!( + header_value(&captured.head, A2A_EXTENSIONS_HEADER).as_deref(), + Some(TASK_FEEDBACK_EXTENSION_URI) + ); + let body: Value = serde_json::from_slice(&captured.body).unwrap(); + assert_eq!( + body, + serde_json::json!({"rating": "down", "comment": "meh"}) + ); + } + + #[test] + fn an_over_long_comment_is_refused_before_any_request() { + // Unreachable base URL: a request would fail with a connect error, not + // with the message asserted here. + let client = Client::new("http://127.0.0.1:1", "sk-test").unwrap(); + let err = submit_task_feedback( + &client, + "ws-1", + "agt-1", + "run-1", + &TaskFeedbackRequest { + rating: Rating::Up, + comment: Some("x".repeat(FEEDBACK_COMMENT_MAX_CHARS + 1)), + }, + ) + .expect_err("too long"); + assert!(err.to_string().contains("2001 characters"), "{err}"); + } + + #[test] + fn cancel_of_a_finished_task_surfaces_the_a2a_reason() { + // The production response, verbatim: an envelope whose `message` is + // the A2A error document as a string. + let body = r#"{"success":false,"message":"{\"error\":{\"code\":400,\"status\":\"FAILED_PRECONDITION\",\"message\":\"Task cannot be canceled - current state: 3\",\"details\":[{\"@type\":\"type.googleapis.com/google.rpc.ErrorInfo\",\"reason\":\"TASK_NOT_CANCELABLE\",\"domain\":\"a2a-protocol.org\",\"metadata\":{}}]}}","error_code":"INVALID_ARGUMENT"}"#; + let (base_url, _server) = one_shot_server(format!( + "HTTP/1.1 400 Bad Request\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{body}", + body.len() + )); + let client = Client::new(base_url, "sk-test").unwrap(); + + let err = cancel_task(&client, "ws-1", "agt-1", "run-1").expect_err("not cancelable"); + let text = err.to_string(); + assert!( + text.starts_with("Task cannot be canceled - current state: 3 (TASK_NOT_CANCELABLE)"), + "{text}" + ); + assert!( + matches!(&err, Error::Api { code: Some(code), .. } if code == "INVALID_ARGUMENT"), + "{err:?}" + ); + } +} diff --git a/crates/core/src/api/agents/a2a/types.rs b/crates/core/src/api/agents/a2a/types.rs new file mode 100644 index 0000000..4429d91 --- /dev/null +++ b/crates/core/src/api/agents/a2a/types.rs @@ -0,0 +1,251 @@ +//! Request shapes for the A2A v1.0 `SendMessage` operation. +//! +//! Field names are camelCase on the wire, as A2A spells them, unlike the +//! snake_case used by the rest of the MemoryLake API. Only the fields the CLI +//! sets are modelled; message parts are open JSON so the caller can send any +//! part kind the protocol defines. + +use serde::Serialize; +use serde_json::{Map, Value}; + +/// The `role` of a message sent by the caller. +pub const ROLE_USER: &str = "ROLE_USER"; + +/// A task the agent is still working on. +pub const TASK_STATE_WORKING: &str = "TASK_STATE_WORKING"; +/// A task that finished and produced its result. +pub const TASK_STATE_COMPLETED: &str = "TASK_STATE_COMPLETED"; +/// A task that stopped with an error. +pub const TASK_STATE_FAILED: &str = "TASK_STATE_FAILED"; +/// A task that was cancelled. +pub const TASK_STATE_CANCELED: &str = "TASK_STATE_CANCELED"; +/// A task the agent refused. +pub const TASK_STATE_REJECTED: &str = "TASK_STATE_REJECTED"; +/// A task waiting for the caller to say more; reply with its `taskId`. +pub const TASK_STATE_INPUT_REQUIRED: &str = "TASK_STATE_INPUT_REQUIRED"; + +/// Whether a task in `state` will change no further on its own. +/// +/// `TASK_STATE_INPUT_REQUIRED` is *not* terminal: the task resumes when the +/// caller replies to it. +pub fn is_terminal_state(state: &str) -> bool { + matches!( + state, + TASK_STATE_COMPLETED | TASK_STATE_FAILED | TASK_STATE_CANCELED | TASK_STATE_REJECTED + ) +} + +/// A text part, the one part kind the CLI builds itself. +pub fn text_part(text: impl Into) -> Value { + let mut part = Map::new(); + part.insert("text".into(), Value::String(text.into())); + Value::Object(part) +} + +/// Body of `POST .../a2a/message:send` and `message:stream`. +#[derive(Debug, Clone, PartialEq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct SendMessageRequest { + /// The message to deliver. + pub message: Message, + /// How the server should answer. + #[serde(skip_serializing_if = "Option::is_none")] + pub configuration: Option, + /// Request metadata; MemoryLake reads its own extension from here. + #[serde(skip_serializing_if = "Option::is_none")] + pub metadata: Option, +} + +/// One A2A message. +#[derive(Debug, Clone, PartialEq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct Message { + /// Who is speaking; the CLI always sends [`ROLE_USER`]. + pub role: String, + /// Caller-chosen id, unique within the context. + pub message_id: String, + /// Content parts, e.g. `{"text": "..."}`. + pub parts: Vec, + /// Conversation thread to continue. Omitted starts a new one. + #[serde(skip_serializing_if = "Option::is_none")] + pub context_id: Option, + /// Task to continue, for answering a `TASK_STATE_INPUT_REQUIRED` task. + #[serde(skip_serializing_if = "Option::is_none")] + pub task_id: Option, +} + +/// `configuration` of a send request. +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct SendConfiguration { + /// Answer as soon as the task exists instead of when it finishes. + /// + /// The server blocks by default (measured: a short answer returns + /// `TASK_STATE_COMPLETED` with its artifacts in a few seconds); `true` + /// returns the task in `TASK_STATE_WORKING` for the caller to poll. + #[serde(skip_serializing_if = "Option::is_none")] + pub return_immediately: Option, + /// How many history messages to include in the returned task. + #[serde(skip_serializing_if = "Option::is_none")] + pub history_length: Option, +} + +impl SendConfiguration { + /// Whether nothing was set, so the whole object can be left out. + pub fn is_empty(&self) -> bool { + *self == Self::default() + } +} + +/// `metadata` of a send request. +#[derive(Debug, Clone, PartialEq, Serialize)] +pub struct SendMetadata { + /// MemoryLake's extension: which actor speaks, which projects are in + /// scope, whether to remember the exchange. + pub memorylake: MemorylakeExtension, +} + +/// `metadata.memorylake` of a send request. +/// +/// Everything is optional; the server falls back to the agent's own defaults +/// for anything left out. `extra` carries keys the CLI does not model +/// (`overrides`, `subagentMapping`, …) verbatim. +#[derive(Debug, Clone, Default, PartialEq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct MemorylakeExtension { + /// Actor the message is attributed to. Need not be bound to the + /// workspace: production accepts the caller's own actor as is. + #[serde(skip_serializing_if = "Option::is_none")] + pub actor_id: Option, + /// Project the agent may read from and write memories to. + #[serde(skip_serializing_if = "Option::is_none")] + pub read_write_project_id: Option, + /// Further projects the agent may read but not write. + #[serde(skip_serializing_if = "Vec::is_empty")] + pub read_only_project_ids: Vec, + /// Do not extract memories from this exchange. + #[serde(skip_serializing_if = "Option::is_none")] + pub skip_memory: Option, + /// Unmodelled keys, merged into the object as given. + #[serde(flatten)] + pub extra: Map, +} + +impl MemorylakeExtension { + /// Whether nothing was set, so `metadata` can be left out entirely. + pub fn is_empty(&self) -> bool { + self.actor_id.is_none() + && self.read_write_project_id.is_none() + && self.read_only_project_ids.is_empty() + && self.skip_memory.is_none() + && self.extra.is_empty() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn minimal() -> SendMessageRequest { + SendMessageRequest { + message: Message { + role: ROLE_USER.into(), + message_id: "m-1".into(), + parts: vec![text_part("hi")], + context_id: None, + task_id: None, + }, + configuration: None, + metadata: None, + } + } + + #[test] + fn a_minimal_request_sends_only_the_message() { + assert_eq!( + serde_json::to_value(minimal()).unwrap(), + json!({"message": {"role": "ROLE_USER", "messageId": "m-1", "parts": [{"text": "hi"}]}}) + ); + } + + #[test] + fn field_names_are_camel_case_on_the_wire() { + let mut request = minimal(); + request.message.context_id = Some("ctx".into()); + request.message.task_id = Some("run-1".into()); + request.configuration = Some(SendConfiguration { + return_immediately: Some(true), + history_length: Some(3), + }); + request.metadata = Some(SendMetadata { + memorylake: MemorylakeExtension { + actor_id: Some("act-1".into()), + read_write_project_id: Some("proj-1".into()), + read_only_project_ids: vec!["proj-2".into()], + skip_memory: Some(true), + extra: Map::new(), + }, + }); + assert_eq!( + serde_json::to_value(request).unwrap(), + json!({ + "message": { + "role": "ROLE_USER", "messageId": "m-1", "parts": [{"text": "hi"}], + "contextId": "ctx", "taskId": "run-1" + }, + "configuration": {"returnImmediately": true, "historyLength": 3}, + "metadata": {"memorylake": { + "actorId": "act-1", + "readWriteProjectId": "proj-1", + "readOnlyProjectIds": ["proj-2"], + "skipMemory": true + }} + }) + ); + } + + #[test] + fn extra_extension_keys_are_flattened_into_the_object() { + let mut extra = Map::new(); + extra.insert("overrides".into(), json!({"model": "x"})); + let ext = MemorylakeExtension { + extra, + ..Default::default() + }; + assert!(!ext.is_empty()); + assert_eq!( + serde_json::to_value(ext).unwrap(), + json!({"overrides": {"model": "x"}}) + ); + } + + #[test] + fn empty_extension_and_configuration_report_so() { + assert!(MemorylakeExtension::default().is_empty()); + assert!(SendConfiguration::default().is_empty()); + assert!( + !SendConfiguration { + return_immediately: Some(false), + ..Default::default() + } + .is_empty(), + "an explicit false is still a setting" + ); + } + + #[test] + fn only_finished_states_are_terminal() { + for state in [ + TASK_STATE_COMPLETED, + TASK_STATE_FAILED, + TASK_STATE_CANCELED, + TASK_STATE_REJECTED, + ] { + assert!(is_terminal_state(state), "{state}"); + } + assert!(!is_terminal_state(TASK_STATE_WORKING)); + assert!(!is_terminal_state(TASK_STATE_INPUT_REQUIRED)); + assert!(!is_terminal_state("TASK_STATE_UNSPECIFIED")); + } +} diff --git a/crates/core/src/api/agents/mod.rs b/crates/core/src/api/agents/mod.rs index 5be51c9..e3e6fe5 100644 --- a/crates/core/src/api/agents/mod.rs +++ b/crates/core/src/api/agents/mod.rs @@ -4,7 +4,10 @@ //! [`update_agent`]. Agent *configuration* (model, policies, prompt, …) is //! immutable: changing it means creating a new version with //! [`create_agent_version`]. +//! +//! Talking *to* a bound agent goes through the A2A protocol in [`a2a`]. +pub mod a2a; mod bind; mod create; mod create_version; diff --git a/crates/core/src/client.rs b/crates/core/src/client.rs index acd1db8..3b1d2f7 100644 --- a/crates/core/src/client.rs +++ b/crates/core/src/client.rs @@ -5,7 +5,7 @@ use std::collections::BTreeMap; use reqwest::blocking::{Body, Client as HttpClient}; use reqwest::header::{ - AUTHORIZATION, CONTENT_DISPOSITION, CONTENT_TYPE, ETAG, HeaderMap, HeaderValue, + ACCEPT, AUTHORIZATION, CONTENT_DISPOSITION, CONTENT_TYPE, ETAG, HeaderMap, HeaderValue, }; use reqwest::{StatusCode, Url}; use serde::Serialize; @@ -13,6 +13,7 @@ use serde::de::DeserializeOwned; use serde_json::Value; use crate::error::{Error, Result}; +use crate::sse::EventStream; /// What a completed download reported about itself. #[derive(Debug, Clone, PartialEq, Eq)] @@ -302,6 +303,98 @@ impl Client { }) } + /// Perform a GET against an endpoint whose success body is a bare JSON + /// document rather than a MemoryLake envelope. + /// + /// The A2A endpoints speak the A2A protocol's own shapes on success but + /// still answer errors as a MemoryLake envelope, so a non-2xx response is + /// handed to the envelope decoder for a message consistent with the rest + /// of the API, and a 2xx body is returned as-is. + pub fn get_json_with_headers( + &self, + path: &str, + query: &[(&str, String)], + headers: &[(&str, &str)], + ) -> Result { + let url = self.url(path); + let mut builder = apply_headers(self.http.get(&url), headers).headers(self.auth_headers()?); + for (key, value) in query { + builder = builder.query(&[(key, value)]); + } + decode_bare_json(self.execute(builder.build()?)?) + } + + /// [`Self::get_json_with_headers`] for a POST with a JSON body. + pub fn post_json_with_headers( + &self, + path: &str, + body: &B, + headers: &[(&str, &str)], + ) -> Result + where + B: Serialize, + { + let url = self.url(path); + let request = apply_headers(self.http.post(&url), headers) + .headers(self.auth_headers()?) + .json(body) + .build()?; + decode_bare_json(self.execute(request)?) + } + + /// Perform a POST whose success response is a `text/event-stream`, and + /// return the events as they arrive. + /// + /// Errors follow the same rule as [`Self::get_json_with_headers`]: a + /// non-2xx status, or a 2xx that is not an event stream, is decoded as a + /// MemoryLake envelope. The latter has been observed on the A2A stream + /// endpoint, which answers a permission failure with a plain JSON body. + pub fn post_event_stream( + &self, + path: &str, + body: &B, + headers: &[(&str, &str)], + ) -> Result>> + where + B: Serialize, + { + let url = self.url(path); + let request = apply_headers(self.http.post(&url), headers) + .headers(self.auth_headers()?) + .header(ACCEPT, "text/event-stream") + .json(body) + .build()?; + let response = self.execute(request)?; + + let status = response.status(); + let content_type = response + .headers() + .get(CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_ascii_lowercase(); + tracing::trace!(status = status.as_u16(), url = %url, content_type, "stream response"); + + if !status.is_success() || !content_type.starts_with("text/event-stream") { + return match validate_envelope(response) { + Err(err) => Err(err), + Ok(_) => Err(Error::Api { + message: format!( + "expected a text/event-stream from {url}, got {} with status {status}", + if content_type.is_empty() { + "no content type" + } else { + content_type.as_str() + } + ), + code: None, + }), + }; + } + + Ok(EventStream::new(std::io::BufReader::new(response))) + } + /// Perform a DELETE whose successful response carries no usable payload. /// /// Prefer this over `delete_data::<()>` for endpoints documented to answer @@ -556,6 +649,15 @@ fn validate_envelope(response: reqwest::blocking::Response) -> Result<(Value, Re "HTTP response" ); + validate_envelope_body(status, url, body) +} + +/// [`validate_envelope`] on a body that has already been read. +fn validate_envelope_body( + status: StatusCode, + url: Url, + body: String, +) -> Result<(Value, ResponseContext)> { if let Ok(auth_err) = serde_json::from_str::(&body) { let server_msg = auth_err .error @@ -632,6 +734,50 @@ where }) } +/// Return a 2xx body as JSON, or decode a non-2xx body as an error envelope. +/// +/// For endpoints that do not wrap successful responses. `validate_envelope` +/// always errors on a non-success status; the other arm exists so that +/// invariant cannot turn into a panic. +fn decode_bare_json(response: reqwest::blocking::Response) -> Result { + let status = response.status(); + if !status.is_success() { + return match validate_envelope(response) { + Err(err) => Err(err), + Ok(_) => Err(Error::Api { + message: format!("request failed with status {status}"), + code: None, + }), + }; + } + + let url = response.url().clone(); + let body = response.text().map_err(Error::from)?; + tracing::trace!(status = status.as_u16(), url = %url, body = %body, "HTTP response"); + + let value: Value = serde_json::from_str(&body).map_err(|err| Error::Api { + message: format!( + "unexpected response from {url} (expected JSON; error: {err})\n{}", + format_http_response(status, &body) + ), + code: None, + })?; + + // A failed envelope behind a 2xx must not pass for a protocol document: + // the endpoints this serves never emit a `success` key of their own. + if value.get("success") == Some(&Value::Bool(false)) { + return match validate_envelope_body(status, url, body) { + Err(err) => Err(err), + Ok(_) => Err(Error::Api { + message: "request failed".into(), + code: None, + }), + }; + } + + Ok(value) +} + fn format_http_response(status: StatusCode, body: &str) -> String { const MAX_BODY: usize = 2_048; diff --git a/crates/core/src/lib.rs b/crates/core/src/lib.rs index 388cf9b..023f873 100644 --- a/crates/core/src/lib.rs +++ b/crates/core/src/lib.rs @@ -5,6 +5,7 @@ pub mod client; pub mod config; pub mod credentials; pub mod error; +pub mod sse; #[cfg(test)] mod test_support; diff --git a/crates/core/src/sse.rs b/crates/core/src/sse.rs new file mode 100644 index 0000000..08ec495 --- /dev/null +++ b/crates/core/src/sse.rs @@ -0,0 +1,172 @@ +//! Server-sent events (`text/event-stream`) parsing. +//! +//! The A2A streaming endpoints answer with one JSON document per event. Only +//! the `data:` field is meaningful there; `event:`, `id:` and `retry:` are +//! accepted and ignored so a server that starts sending them cannot break the +//! stream. The parser is a plain iterator over a `BufRead`, so it is driven by +//! the blocking HTTP response directly and tested with an in-memory cursor. + +use std::io::BufRead; + +use serde_json::Value; + +use crate::error::{Error, Result}; + +/// Iterator over the JSON payloads of a `text/event-stream` body. +/// +/// Each item is one event's `data` decoded as JSON. Events with no `data` +/// lines (keep-alive comments, bare `event:` lines) are skipped rather than +/// reported, since they carry nothing the caller can act on. +#[derive(Debug)] +pub struct EventStream { + reader: R, + /// Set once the body ended or a read failed, so a caller that keeps + /// polling after `None` gets `None` again rather than a fresh read. + done: bool, +} + +impl EventStream { + /// Wrap a body that is already positioned at the first event. + pub fn new(reader: R) -> Self { + Self { + reader, + done: false, + } + } + + /// Read lines up to the next blank line and collect the `data` field. + /// + /// Returns `Ok(None)` at end of body. A body that ends without a trailing + /// blank line still yields its final event: A2A servers close the + /// connection right after the last frame, and the event is complete + /// either way. + fn next_data(&mut self) -> Result> { + let mut data: Vec = Vec::new(); + let mut line = String::new(); + loop { + line.clear(); + let read = self + .reader + .read_line(&mut line) + .map_err(|source| Error::Io { + action: "read event stream", + path: std::path::PathBuf::from(""), + source, + })?; + if read == 0 { + self.done = true; + return Ok((!data.is_empty()).then(|| data.join("\n"))); + } + + let line = line.trim_end_matches(['\r', '\n']); + if line.is_empty() { + if data.is_empty() { + // Blank line between events with nothing pending: a + // keep-alive or a stray separator. Keep reading. + continue; + } + return Ok(Some(data.join("\n"))); + } + if line.starts_with(':') { + // Comment line, used as a keep-alive by some servers. + continue; + } + let (field, value) = match line.split_once(':') { + Some((field, value)) => (field, value.strip_prefix(' ').unwrap_or(value)), + None => (line, ""), + }; + if field == "data" { + data.push(value.to_string()); + } + } + } +} + +impl Iterator for EventStream { + type Item = Result; + + fn next(&mut self) -> Option { + if self.done { + return None; + } + let data = match self.next_data() { + Err(err) => { + self.done = true; + return Some(Err(err)); + } + Ok(None) => return None, + Ok(Some(data)) => data, + }; + let parsed = serde_json::from_str(&data).map_err(|source| Error::Api { + message: format!("event stream sent a frame that is not JSON: {source}\n{data}"), + code: None, + }); + if parsed.is_err() { + // A garbled frame means the rest cannot be trusted to line up + // either; stop rather than resynchronise. + self.done = true; + } + Some(parsed) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::io::Cursor; + + fn events(body: &str) -> Vec { + EventStream::new(Cursor::new(body.as_bytes())) + .collect::>>() + .expect("parse events") + } + + #[test] + fn splits_frames_on_blank_lines() { + // Exactly the shape production sends: `data:` with no space, one + // frame per blank-line-separated block. + let body = "data:{\"a\":1}\n\ndata:{\"b\":2}\n\n"; + assert_eq!( + events(body), + vec![serde_json::json!({"a": 1}), serde_json::json!({"b": 2})] + ); + } + + #[test] + fn accepts_a_space_after_the_colon_and_crlf_line_endings() { + let body = "data: {\"a\":1}\r\n\r\ndata: {\"b\":2}\r\n\r\n"; + assert_eq!(events(body).len(), 2); + } + + #[test] + fn joins_multi_line_data_with_newlines() { + let body = "data:{\"text\":\ndata:\"x\"}\n\n"; + assert_eq!(events(body), vec![serde_json::json!({"text": "x"})]); + } + + #[test] + fn ignores_comments_ids_and_event_names() { + let body = ": keep-alive\nevent: update\nid: 7\nretry: 100\ndata:{\"a\":1}\n\n: ping\n\n"; + assert_eq!(events(body), vec![serde_json::json!({"a": 1})]); + } + + #[test] + fn a_final_frame_without_a_trailing_blank_line_is_still_delivered() { + let body = "data:{\"a\":1}\n\ndata:{\"b\":2}"; + assert_eq!(events(body).len(), 2); + } + + #[test] + fn an_empty_body_yields_nothing() { + assert!(events("").is_empty()); + assert!(events("\n\n: only comments\n\n").is_empty()); + } + + #[test] + fn a_non_json_frame_is_an_error_that_ends_the_stream() { + let mut stream = EventStream::new(Cursor::new(b"data:not json\n\ndata:{\"a\":1}\n\n")); + let err = stream.next().expect("one item").expect_err("not json"); + assert!(err.to_string().contains("not JSON"), "{err}"); + assert!(stream.next().is_none(), "stream must stop after an error"); + } +}