gemini_adk_rs/text/
map_over.rs1use 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
11pub 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 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 pub fn with_middleware_chain(mut self, chain: MiddlewareChain) -> Self {
44 self.middleware = chain;
45 self
46 }
47
48 pub fn item_key(mut self, key: impl Into<String>) -> Self {
50 self.item_key = key.into();
51 self
52 }
53
54 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}