1use std::collections::HashMap;
2
3use bamboo_a2a::types::{A2ARole, PartContentWire, StreamResponse, TaskState, TaskStatus};
4use bamboo_agent_core::{AgentEvent, TokenUsage};
5
6pub struct A2AMappedEvents {
8 pub events: Vec<AgentEvent>,
9 pub metadata_updates: HashMap<String, String>,
10}
11
12#[derive(Default)]
14pub struct A2AEventMapper {
15 terminal_sent: bool,
16 latest_task_id: Option<String>,
17 context_id: Option<String>,
18 final_text: String,
19}
20
21impl A2AEventMapper {
22 pub fn new() -> Self {
23 Self::default()
24 }
25
26 pub fn latest_task_id(&self) -> Option<&str> {
27 self.latest_task_id.as_deref()
28 }
29
30 pub fn context_id(&self) -> Option<&str> {
31 self.context_id.as_deref()
32 }
33
34 pub fn is_terminal(&self) -> bool {
35 self.terminal_sent
36 }
37
38 pub fn final_text(&self) -> &str {
39 &self.final_text
40 }
41
42 pub fn map_stream_response(&mut self, response: StreamResponse) -> A2AMappedEvents {
44 let mut events = Vec::new();
45 let mut metadata = HashMap::new();
46
47 if let Some(task) = response.task {
48 self.latest_task_id = Some(task.id.clone());
49 if let Some(ctx) = task.context_id.clone() {
50 self.context_id = Some(ctx);
51 }
52 metadata.insert("a2a.latest_task_id".to_string(), task.id.clone());
53 if let Some(ctx) = &task.context_id {
54 metadata.insert("a2a.context_id".to_string(), ctx.clone());
55 }
56 metadata.insert(
57 "a2a.last_state".to_string(),
58 task.status.state.as_proto_str().to_string(),
59 );
60 events.extend(self.map_status(&task.id, task.context_id.as_deref(), task.status));
61 }
62
63 if let Some(message) = response.message {
64 if message.role == A2ARole::Agent {
65 let text = text_from_parts(&message.parts);
66 if !text.is_empty() {
67 self.final_text.push_str(&text);
68 events.push(AgentEvent::Token { content: text });
69 }
70 }
71 }
72
73 if let Some(update) = response.status_update {
74 self.latest_task_id = Some(update.task_id.clone());
75 self.context_id = Some(update.context_id.clone());
76 metadata.insert("a2a.latest_task_id".to_string(), update.task_id.clone());
77 metadata.insert("a2a.context_id".to_string(), update.context_id.clone());
78 metadata.insert(
79 "a2a.last_state".to_string(),
80 update.status.state.as_proto_str().to_string(),
81 );
82 events.extend(self.map_status(
83 &update.task_id,
84 Some(&update.context_id),
85 update.status,
86 ));
87 }
88
89 if let Some(update) = response.artifact_update {
90 let preview =
91 handle_artifact_update(&update.artifact, update.append, update.last_chunk);
92 if !preview.is_empty() {
93 events.push(AgentEvent::Token {
94 content: preview.clone(),
95 });
96 self.final_text.push_str(&preview);
97 }
98 metadata.insert(
99 "a2a.last_artifacts_summary".to_string(),
100 serde_json::json!({
101 "artifact_id": update.artifact.artifact_id,
102 "name": update.artifact.name,
103 "append": update.append,
104 "last_chunk": update.last_chunk,
105 })
106 .to_string(),
107 );
108 }
109
110 A2AMappedEvents {
111 events,
112 metadata_updates: metadata,
113 }
114 }
115
116 fn map_status(
117 &mut self,
118 _task_id: &str,
119 _context_id: Option<&str>,
120 status: TaskStatus,
121 ) -> Vec<AgentEvent> {
122 let mut events = Vec::new();
123
124 match &status.state {
125 TaskState::Submitted => {
126 }
128 TaskState::Working => {
129 if let Some(msg) = &status.message {
130 let text = text_from_parts(&msg.parts);
131 if !text.is_empty() {
132 self.final_text.push_str(&text);
133 events.push(AgentEvent::Token { content: text });
134 }
135 }
136 }
137 TaskState::InputRequired => {
138 let question = question_from_status(&status);
139 events.push(AgentEvent::NeedClarification {
140 question,
141 options: None,
142 tool_call_id: None,
143 tool_name: None,
144 allow_custom: true,
145 source: Some(bamboo_agent_core::PendingQuestionSource::ExternalAgent),
146 });
147 }
148 TaskState::AuthRequired => {
149 let question = question_from_status(&status);
150 events.push(AgentEvent::NeedClarification {
151 question,
152 options: None,
153 tool_call_id: None,
154 tool_name: None,
155 allow_custom: true,
156 source: Some(bamboo_agent_core::PendingQuestionSource::ExternalAgent),
157 });
158 }
159 TaskState::Completed => {
160 self.terminal_sent = true;
161 if let Some(msg) = &status.message {
162 let text = text_from_parts(&msg.parts);
163 if !text.is_empty() {
164 self.final_text.push_str(&text);
165 events.push(AgentEvent::Token { content: text });
166 }
167 }
168 events.push(AgentEvent::Complete {
169 usage: TokenUsage::default(),
170 });
171 }
172 TaskState::Failed => {
173 self.terminal_sent = true;
174 let error_msg = status
175 .message
176 .as_ref()
177 .map(|m| text_from_parts(&m.parts))
178 .filter(|s| !s.is_empty())
179 .unwrap_or_else(|| "External agent reported failure".to_string());
180 events.push(AgentEvent::Error { message: error_msg });
181 }
182 TaskState::Canceled => {
183 self.terminal_sent = true;
184 events.push(AgentEvent::Error {
185 message: "External agent task was cancelled".to_string(),
186 });
187 }
188 TaskState::Rejected => {
189 self.terminal_sent = true;
190 events.push(AgentEvent::Error {
191 message: "External agent rejected the task".to_string(),
192 });
193 }
194 TaskState::Unspecified => {}
195 }
196
197 events
198 }
199}
200
201pub fn text_from_parts(parts: &[bamboo_a2a::types::Part]) -> String {
203 parts
204 .iter()
205 .filter_map(|part| match &part.content {
206 PartContentWire::Text { text } => Some(text.as_str()),
207 PartContentWire::Data { data } => data.get("summary").and_then(|v| v.as_str()),
208 _ => None,
209 })
210 .collect::<Vec<_>>()
211 .join("\n")
212}
213
214fn question_from_status(status: &TaskStatus) -> String {
216 status
217 .message
218 .as_ref()
219 .map(|m| text_from_parts(&m.parts))
220 .filter(|s| !s.trim().is_empty())
221 .unwrap_or_else(|| match status.state {
222 TaskState::InputRequired => "External agent requires additional input.".to_string(),
223 TaskState::AuthRequired => {
224 "External agent requires authentication or authorization.".to_string()
225 }
226 _ => format!("External agent state: {:?}", status.state),
227 })
228}
229
230fn handle_artifact_update(
232 artifact: &bamboo_a2a::types::Artifact,
233 _append: bool,
234 _last_chunk: bool,
235) -> String {
236 let text = text_from_parts(&artifact.parts);
237 if text.is_empty() {
238 if let Some(name) = &artifact.name {
239 format!("[Artifact: {}]", name)
240 } else {
241 format!("[Artifact: {}]", artifact.artifact_id)
242 }
243 } else {
244 let header = artifact
245 .name
246 .as_ref()
247 .map(|n| format!("--- Artifact: {} ---\n", n))
248 .unwrap_or_default();
249 format!("{}{}", header, text)
250 }
251}
252
253#[cfg(test)]
254mod tests {
255 use super::*;
256 use bamboo_a2a::types::{A2ARole, Message, Part, Task, TaskStatus, TaskStatusUpdateEvent};
257
258 #[test]
259 fn a2a_message_text_maps_to_token() {
260 let mut mapper = A2AEventMapper::new();
261 let response = StreamResponse {
262 task: None,
263 message: Some(Message {
264 message_id: "m1".to_string(),
265 context_id: None,
266 task_id: None,
267 role: A2ARole::Agent,
268 parts: vec![Part {
269 content: PartContentWire::text("hello world"),
270 metadata: None,
271 filename: None,
272 media_type: Some("text/plain".to_string()),
273 }],
274 metadata: None,
275 extensions: vec![],
276 reference_task_ids: vec![],
277 }),
278 status_update: None,
279 artifact_update: None,
280 };
281 let mapped = mapper.map_stream_response(response);
282 assert_eq!(mapped.events.len(), 1);
283 match &mapped.events[0] {
284 AgentEvent::Token { content } => assert_eq!(content, "hello world"),
285 other => panic!("expected Token, got {:?}", other),
286 }
287 }
288
289 #[test]
290 fn a2a_completed_status_maps_to_complete_and_metadata() {
291 let mut mapper = A2AEventMapper::new();
292 let response = StreamResponse {
293 task: Some(Task {
294 id: "task-1".to_string(),
295 context_id: Some("ctx-1".to_string()),
296 status: TaskStatus {
297 state: TaskState::Completed,
298 message: None,
299 timestamp: None,
300 },
301 artifacts: vec![],
302 history: vec![],
303 metadata: None,
304 }),
305 message: None,
306 status_update: None,
307 artifact_update: None,
308 };
309 let mapped = mapper.map_stream_response(response);
310 assert!(mapper.is_terminal());
311 assert_eq!(
312 mapped.metadata_updates.get("a2a.latest_task_id"),
313 Some(&"task-1".to_string())
314 );
315 assert_eq!(
316 mapped.metadata_updates.get("a2a.context_id"),
317 Some(&"ctx-1".to_string())
318 );
319 assert_eq!(
320 mapped.metadata_updates.get("a2a.last_state"),
321 Some(&"TASK_STATE_COMPLETED".to_string())
322 );
323 match &mapped.events[0] {
324 AgentEvent::Complete { .. } => {}
325 other => panic!("expected Complete, got {:?}", other),
326 }
327 }
328
329 #[test]
330 fn a2a_failed_status_maps_to_error() {
331 let mut mapper = A2AEventMapper::new();
332 let response = StreamResponse {
333 task: None,
334 message: None,
335 status_update: Some(TaskStatusUpdateEvent {
336 task_id: "task-1".to_string(),
337 context_id: "ctx-1".to_string(),
338 status: TaskStatus {
339 state: TaskState::Failed,
340 message: Some(Message {
341 message_id: "m1".to_string(),
342 context_id: None,
343 task_id: None,
344 role: A2ARole::Agent,
345 parts: vec![Part {
346 content: PartContentWire::text("Something went wrong"),
347 metadata: None,
348 filename: None,
349 media_type: None,
350 }],
351 metadata: None,
352 extensions: vec![],
353 reference_task_ids: vec![],
354 }),
355 timestamp: None,
356 },
357 metadata: None,
358 }),
359 artifact_update: None,
360 };
361 let mapped = mapper.map_stream_response(response);
362 assert!(mapper.is_terminal());
363 match &mapped.events[0] {
364 AgentEvent::Error { message } => assert_eq!(message, "Something went wrong"),
365 other => panic!("expected Error, got {:?}", other),
366 }
367 }
368
369 #[test]
370 fn a2a_input_required_maps_to_need_clarification() {
371 let mut mapper = A2AEventMapper::new();
372 let response = StreamResponse {
373 task: None,
374 message: None,
375 status_update: Some(TaskStatusUpdateEvent {
376 task_id: "task-1".to_string(),
377 context_id: "ctx-1".to_string(),
378 status: TaskStatus {
379 state: TaskState::InputRequired,
380 message: Some(Message {
381 message_id: "m1".to_string(),
382 context_id: None,
383 task_id: None,
384 role: A2ARole::Agent,
385 parts: vec![Part {
386 content: PartContentWire::text("What is your API key?"),
387 metadata: None,
388 filename: None,
389 media_type: None,
390 }],
391 metadata: None,
392 extensions: vec![],
393 reference_task_ids: vec![],
394 }),
395 timestamp: None,
396 },
397 metadata: None,
398 }),
399 artifact_update: None,
400 };
401 let mapped = mapper.map_stream_response(response);
402 assert!(!mapper.is_terminal());
403 match &mapped.events[0] {
404 AgentEvent::NeedClarification { question, .. } => {
405 assert_eq!(question, "What is your API key?");
406 }
407 other => panic!("expected NeedClarification, got {:?}", other),
408 }
409 }
410}