Skip to main content

bamboo_engine/external_agents/
mapping.rs

1use std::collections::HashMap;
2
3use bamboo_a2a::types::{A2ARole, PartContentWire, StreamResponse, TaskState, TaskStatus};
4use bamboo_agent_core::{AgentEvent, TokenUsage};
5
6/// Events and metadata updates produced by mapping a single A2A StreamResponse.
7pub struct A2AMappedEvents {
8    pub events: Vec<AgentEvent>,
9    pub metadata_updates: HashMap<String, String>,
10}
11
12/// Stateful mapper that tracks the latest A2A task state across a stream.
13#[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    /// Map a single A2A StreamResponse to Bamboo AgentEvents and metadata.
43    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                // Optionally emit a brief status token
127            }
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
201/// Extract plain text from a slice of Parts.
202pub 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
214/// Build a human-readable question from a TaskStatus that requires input/auth.
215fn 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
230/// Build a preview string from an artifact update.
231fn 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}