gemini_adk_rs/text/
map_over.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4
5use super::TextAgent;
6use crate::context::AgentEvent;
7use crate::error::AgentError;
8use crate::middleware::MiddlewareChain;
9use crate::state::State;
10
11/// Iterates a single agent over each item in a state list.
12/// Reads `state[list_key]`, runs agent per item (setting `state[item_key]`),
13/// collects results into `state[output_key]`.
14pub struct MapOverTextAgent {
15    name: String,
16    agent: Arc<dyn TextAgent>,
17    list_key: String,
18    item_key: String,
19    output_key: String,
20    middleware: MiddlewareChain,
21}
22
23impl MapOverTextAgent {
24    /// Create a new map-over agent that iterates over a list in state.
25    pub fn new(
26        name: impl Into<String>,
27        agent: Arc<dyn TextAgent>,
28        list_key: impl Into<String>,
29    ) -> Self {
30        Self {
31            name: name.into(),
32            agent,
33            list_key: list_key.into(),
34            item_key: "_item".into(),
35            output_key: "_results".into(),
36            middleware: MiddlewareChain::new(),
37        }
38    }
39
40    /// Attach a middleware chain. `AgentEvent::LoopIteration` is emitted
41    /// through it before each item is processed (zero-based), so `on_event`
42    /// observers (e.g. `M::on_loop`) fire per item.
43    pub fn with_middleware_chain(mut self, chain: MiddlewareChain) -> Self {
44        self.middleware = chain;
45        self
46    }
47
48    /// Set the state key for the current item (default: "_item").
49    pub fn item_key(mut self, key: impl Into<String>) -> Self {
50        self.item_key = key.into();
51        self
52    }
53
54    /// Set the state key for the output list (default: "_results").
55    pub fn output_key(mut self, key: impl Into<String>) -> Self {
56        self.output_key = key.into();
57        self
58    }
59}
60
61#[async_trait]
62impl TextAgent for MapOverTextAgent {
63    fn name(&self) -> &str {
64        &self.name
65    }
66
67    async fn run(&self, state: &State) -> Result<String, AgentError> {
68        let items: Vec<serde_json::Value> = state.get(&self.list_key).unwrap_or_default();
69
70        let mut results = Vec::with_capacity(items.len());
71
72        for (iteration, item) in items.iter().enumerate() {
73            let _ = self
74                .middleware
75                .run_on_event(&AgentEvent::LoopIteration {
76                    iteration: iteration as u32,
77                })
78                .await;
79            let _ = state.set(&self.item_key, item);
80            let _ = state.set("input", item.to_string());
81            let result = self.agent.run(state).await?;
82            results.push(result);
83        }
84
85        let _ = state.set(&self.output_key, &results);
86        Ok(results.join("\n"))
87    }
88}