1use std::sync::Arc;
7
8use futures_util::stream::{self, BoxStream, StreamExt};
9
10use crate::artifacts::{ArtifactService, InMemoryArtifactService};
11use crate::error::AgentError;
12use crate::events::{Event, EventActions};
13use crate::memory::{InMemoryMemoryService, MemoryService};
14use crate::plugin::{Plugin, PluginManager};
15use crate::session::{InMemorySessionService, SessionId, SessionService};
16use crate::state::State;
17use crate::text::TextAgent;
18
19#[derive(Debug)]
26pub enum RunEvent {
27 Event(Box<Event>),
30 Error(AgentError),
32}
33
34impl RunEvent {
35 fn event_item(event: Event) -> Self {
37 RunEvent::Event(Box::new(event))
38 }
39
40 pub fn event(&self) -> Option<&Event> {
42 match self {
43 RunEvent::Event(e) => Some(e),
44 RunEvent::Error(_) => None,
45 }
46 }
47}
48
49enum RunStep {
51 Start,
53 Run {
55 session_id: SessionId,
56 state: State,
57 baseline: std::collections::HashMap<String, serde_json::Value>,
58 },
59 Done,
61}
62
63pub struct TextRunner {
68 root_agent: Arc<dyn TextAgent>,
69 session_service: Arc<dyn SessionService>,
70 memory_service: Arc<dyn MemoryService>,
71 artifact_service: Arc<dyn ArtifactService>,
72 plugins: PluginManager,
73 app_name: String,
74}
75
76impl TextRunner {
77 pub fn new(agent: Arc<dyn TextAgent>, app_name: impl Into<String>) -> Self {
79 Self {
80 root_agent: agent,
81 session_service: Arc::new(InMemorySessionService::new()),
82 memory_service: Arc::new(InMemoryMemoryService::new()),
83 artifact_service: Arc::new(InMemoryArtifactService::new()),
84 plugins: PluginManager::new(),
85 app_name: app_name.into(),
86 }
87 }
88
89 pub fn session_service(mut self, svc: Arc<dyn SessionService>) -> Self {
91 self.session_service = svc;
92 self
93 }
94
95 pub fn memory_service(mut self, svc: Arc<dyn MemoryService>) -> Self {
97 self.memory_service = svc;
98 self
99 }
100
101 pub fn artifact_service(mut self, svc: Arc<dyn ArtifactService>) -> Self {
103 self.artifact_service = svc;
104 self
105 }
106
107 pub fn plugin(mut self, p: impl Plugin + 'static) -> Self {
109 self.plugins.add(Arc::new(p));
110 self
111 }
112
113 pub async fn run(
125 &self,
126 prompt: &str,
127 user_id: &str,
128 session_id: Option<&SessionId>,
129 ) -> Result<String, AgentError> {
130 let mut stream = self.run_stream(prompt, user_id, session_id);
131 let mut last_response: Option<String> = None;
132 while let Some(item) = stream.next().await {
133 match item {
134 RunEvent::Error(e) => return Err(e),
135 RunEvent::Event(ev) => {
136 if ev.author == self.root_agent.name() {
138 last_response = ev.content.clone();
139 }
140 }
141 }
142 }
143 Ok(last_response.unwrap_or_default())
144 }
145
146 pub fn run_stream<'a>(
166 &'a self,
167 prompt: &'a str,
168 user_id: &'a str,
169 session_id: Option<&'a SessionId>,
170 ) -> BoxStream<'a, RunEvent> {
171 let prompt = prompt.to_string();
175 let user_id = user_id.to_string();
176 let session_id = session_id.cloned();
177
178 stream::unfold(RunStep::Start, move |step| {
179 let prompt = prompt.clone();
180 let user_id = user_id.clone();
181 let session_id = session_id.clone();
182 async move {
183 match step {
184 RunStep::Start => {
185 let session = match &session_id {
187 Some(id) => match self.session_service.get_session(id).await {
188 Ok(Some(s)) => s,
189 Ok(None) => {
190 return Some((
191 RunEvent::Error(AgentError::Other(format!(
192 "Session not found: {id}"
193 ))),
194 RunStep::Done,
195 ));
196 }
197 Err(e) => {
198 return Some((
199 RunEvent::Error(AgentError::Other(format!(
200 "Session error: {e}"
201 ))),
202 RunStep::Done,
203 ));
204 }
205 },
206 None => match self
207 .session_service
208 .create_session(&self.app_name, &user_id)
209 .await
210 {
211 Ok(s) => s,
212 Err(e) => {
213 return Some((
214 RunEvent::Error(AgentError::Other(format!(
215 "Session create error: {e}"
216 ))),
217 RunStep::Done,
218 ));
219 }
220 },
221 };
222
223 let state = State::new();
225 let prior = match self.session_service.get_events(&session.id).await {
226 Ok(evs) => evs,
227 Err(e) => {
228 return Some((
229 RunEvent::Error(AgentError::Other(format!(
230 "Events error: {e}"
231 ))),
232 RunStep::Done,
233 ));
234 }
235 };
236 for event in &prior {
237 let encoded = event.actions.is_format_marked();
243 for (key, value) in &event.actions.state_delta {
244 if encoded
245 && (key == EventActions::REMOVED_KEYS
246 || key == EventActions::FORMAT)
247 {
248 continue;
249 }
250 let key = if encoded {
251 EventActions::decode_key(key).into_owned()
252 } else {
253 key.clone()
254 };
255 let _ = state.set(key, value.clone());
256 }
257 for key in event.actions.removed_keys() {
262 let _ = state.remove(key);
263 }
264 }
265 let _ = state.set("input", &prompt);
266
267 let baseline = state.to_hashmap();
269
270 let user_event = Event::new("user", Some(prompt.clone()));
272 if let Err(e) = self
273 .session_service
274 .append_event(&session.id, user_event.clone())
275 .await
276 {
277 return Some((
278 RunEvent::Error(AgentError::Other(format!(
279 "Event append error: {e}"
280 ))),
281 RunStep::Done,
282 ));
283 }
284
285 Some((
287 RunEvent::event_item(user_event),
288 RunStep::Run {
289 session_id: session.id,
290 state,
291 baseline,
292 },
293 ))
294 }
295 RunStep::Run {
296 session_id,
297 state,
298 baseline,
299 } => {
300 let result = match self.root_agent.run(&state).await {
302 Ok(r) => r,
303 Err(e) => return Some((RunEvent::Error(e), RunStep::Done)),
304 };
305
306 let after = state.to_hashmap();
308 let mut delta = std::collections::HashMap::new();
309 for (key, value) in &after {
310 if key == "input" {
311 continue;
312 }
313 if baseline.get(key) != Some(value) {
314 delta.insert(
318 EventActions::encode_key(key).into_owned(),
319 value.clone(),
320 );
321 }
322 }
323 let removed: Vec<serde_json::Value> = baseline
330 .keys()
331 .filter(|key| *key != "input" && !after.contains_key(*key))
332 .map(|key| serde_json::Value::String(key.clone()))
333 .collect();
334 if !removed.is_empty() {
335 delta.insert(
336 EventActions::REMOVED_KEYS.to_string(),
337 serde_json::Value::Array(removed),
338 );
339 }
340
341 let mut actions = EventActions {
342 state_delta: delta,
343 ..Default::default()
344 };
345 actions.mark_format();
349 let result_event = Event::new(self.root_agent.name(), Some(result.clone()))
350 .with_actions(actions);
351
352 if let Err(e) = self
354 .session_service
355 .append_event(&session_id, result_event.clone())
356 .await
357 {
358 return Some((
359 RunEvent::Error(AgentError::Other(format!(
360 "Event append error: {e}"
361 ))),
362 RunStep::Done,
363 ));
364 }
365
366 Some((RunEvent::event_item(result_event), RunStep::Done))
367 }
368 RunStep::Done => None,
369 }
370 }
371 })
372 .boxed()
373 }
374
375 pub async fn run_ephemeral(&self, prompt: &str) -> Result<String, AgentError> {
377 let state = State::new();
378 let _ = state.set("input", prompt);
379 self.root_agent.run(&state).await
380 }
381
382 pub fn session_service_ref(&self) -> &dyn SessionService {
384 self.session_service.as_ref()
385 }
386
387 pub fn app_name(&self) -> &str {
389 &self.app_name
390 }
391}
392
393#[cfg(test)]
394mod tests {
395 use super::*;
396 use crate::text::FnTextAgent;
397
398 fn echo_agent() -> Arc<dyn TextAgent> {
399 Arc::new(FnTextAgent::new("echo", |state| {
400 let input: String = state.get("input").unwrap_or_default();
401 Ok(format!("Echo: {input}"))
402 }))
403 }
404
405 #[tokio::test]
406 async fn run_ephemeral() {
407 let runner = TextRunner::new(echo_agent(), "test-app");
408 let result = runner.run_ephemeral("Hello").await.unwrap();
409 assert_eq!(result, "Echo: Hello");
410 }
411
412 #[tokio::test]
413 async fn run_with_session_creates_and_persists() {
414 let runner = TextRunner::new(echo_agent(), "test-app");
415
416 let result = runner.run("Hello", "user-1", None).await.unwrap();
418 assert_eq!(result, "Echo: Hello");
419
420 let sessions = runner
422 .session_service_ref()
423 .list_sessions("test-app", "user-1")
424 .await
425 .unwrap();
426 assert_eq!(sessions.len(), 1);
427
428 let events = runner
430 .session_service_ref()
431 .get_events(&sessions[0].id)
432 .await
433 .unwrap();
434 assert_eq!(events.len(), 2);
435 assert_eq!(events[0].author, "user");
436 assert_eq!(events[1].author, "echo");
437 }
438
439 #[tokio::test]
440 async fn run_resumes_existing_session() {
441 let runner = TextRunner::new(echo_agent(), "test-app");
442
443 let result1 = runner.run("First", "user-1", None).await.unwrap();
445 assert_eq!(result1, "Echo: First");
446
447 let sessions = runner
449 .session_service_ref()
450 .list_sessions("test-app", "user-1")
451 .await
452 .unwrap();
453 let session_id = &sessions[0].id;
454
455 let result2 = runner
457 .run("Second", "user-1", Some(session_id))
458 .await
459 .unwrap();
460 assert_eq!(result2, "Echo: Second");
461
462 let events = runner
464 .session_service_ref()
465 .get_events(session_id)
466 .await
467 .unwrap();
468 assert_eq!(events.len(), 4);
469 }
470
471 #[tokio::test]
472 async fn run_with_nonexistent_session_errors() {
473 let runner = TextRunner::new(echo_agent(), "test-app");
474 let fake_id = SessionId::new();
475 let result = runner.run("Hello", "user-1", Some(&fake_id)).await;
476 assert!(result.is_err());
477 }
478
479 #[tokio::test]
480 async fn custom_session_service() {
481 let custom_svc = Arc::new(InMemorySessionService::new());
482 let runner = TextRunner::new(echo_agent(), "app").session_service(custom_svc.clone());
483
484 runner.run("Hi", "u1", None).await.unwrap();
485
486 let sessions = custom_svc.list_sessions("app", "u1").await.unwrap();
487 assert_eq!(sessions.len(), 1);
488 }
489
490 fn delta_agent() -> Arc<dyn TextAgent> {
493 Arc::new(FnTextAgent::new("worker", |state| {
494 let input: String = state.get("input").unwrap_or_default();
495 let _ = state.set("turn_count", 1u32);
496 Ok(format!("Handled: {input}"))
497 }))
498 }
499
500 #[tokio::test]
501 async fn run_stream_yields_user_then_final_event() {
502 let runner = TextRunner::new(echo_agent(), "test-app");
503 let mut stream = runner.run_stream("Hello", "user-1", None);
504
505 let mut events = Vec::new();
506 while let Some(item) = stream.next().await {
507 match item {
508 RunEvent::Event(e) => events.push(e),
509 RunEvent::Error(e) => panic!("unexpected error: {e}"),
510 }
511 }
512
513 assert_eq!(events.len(), 2);
515 assert_eq!(events[0].author, "user");
516 assert_eq!(events[0].content.as_deref(), Some("Hello"));
517 assert_eq!(events[1].author, "echo");
518 assert_eq!(events[1].content.as_deref(), Some("Echo: Hello"));
519 }
520
521 #[tokio::test]
522 async fn run_stream_surfaces_state_delta_on_final_event() {
523 let runner = TextRunner::new(delta_agent(), "test-app");
524 let mut stream = runner.run_stream("go", "user-1", None);
525
526 let mut events = Vec::new();
527 while let Some(item) = stream.next().await {
528 if let RunEvent::Event(e) = item {
529 events.push(e);
530 }
531 }
532
533 let final_event = events.last().expect("final event");
534 assert_eq!(final_event.author, "worker");
535 assert_eq!(
536 final_event.actions.state_delta.get("turn_count"),
537 Some(&serde_json::json!(1))
538 );
539 }
540
541 #[tokio::test]
542 async fn removed_keys_are_recorded_and_stay_removed_across_replay() {
543 let agent = Arc::new(FnTextAgent::new("worker", |state| {
547 let input: String = state.get("input").unwrap_or_default();
548 match input.as_str() {
549 "set" => {
550 let _ = state.set("flag", true);
551 }
552 "clear" => {
553 let _ = state.remove("flag");
554 }
555 _ => {}
556 }
557 Ok(format!("saw: {:?}", state.get::<bool>("flag")))
558 }));
559 let runner = TextRunner::new(agent, "test-app");
560
561 runner.run("set", "user-1", None).await.unwrap();
562 let sessions = runner
563 .session_service_ref()
564 .list_sessions("test-app", "user-1")
565 .await
566 .unwrap();
567 let sid = sessions[0].id.clone();
568
569 runner.run("clear", "user-1", Some(&sid)).await.unwrap();
570 let events = runner.session_service_ref().get_events(&sid).await.unwrap();
571 let clear_event = events
572 .iter()
573 .rev()
574 .find(|e| e.author == "worker")
575 .expect("clear run's agent event");
576 assert!(
577 clear_event.actions.removed_keys().any(|k| k == "flag"),
578 "removal persisted under the reserved entry, got {:?}",
579 clear_event.actions.state_delta
580 );
581 assert!(
582 !clear_event.actions.state_delta.contains_key("flag"),
583 "a removal must not also be written as a delta value"
584 );
585
586 let peeked = runner.run("peek", "user-1", Some(&sid)).await.unwrap();
587 assert_eq!(
588 peeked, "saw: None",
589 "removed key did not resurrect on replay"
590 );
591 }
592
593 #[tokio::test]
594 async fn a_state_value_stored_at_the_reserved_key_survives_replay() {
595 let agent = Arc::new(FnTextAgent::new("worker", |state| {
599 let input: String = state.get("input").unwrap_or_default();
600 if input == "store" {
601 let _ = state.set("keep", "kept");
602 let _ = state.set(EventActions::REMOVED_KEYS, vec!["keep"]);
603 }
604 Ok(format!(
605 "{:?}/{:?}",
606 state.get::<Vec<String>>(EventActions::REMOVED_KEYS),
607 state.get::<String>("keep")
608 ))
609 }));
610 let runner = TextRunner::new(agent, "test-app");
611
612 runner.run("store", "user-1", None).await.unwrap();
613 let sessions = runner
614 .session_service_ref()
615 .list_sessions("test-app", "user-1")
616 .await
617 .unwrap();
618 let sid = sessions[0].id.clone();
619
620 let events = runner.session_service_ref().get_events(&sid).await.unwrap();
621 let stored = events
622 .iter()
623 .rev()
624 .find(|e| e.author == "worker")
625 .expect("store run\'s agent event");
626 assert_eq!(
627 stored.actions.removed_keys().count(),
628 0,
629 "a state value must not be mistaken for a removal list"
630 );
631
632 let peeked = runner.run("peek", "user-1", Some(&sid)).await.unwrap();
633 assert_eq!(
634 peeked, "Some([\"keep\"])/Some(\"kept\")",
635 "the value at the reserved key was dropped, or read as a deletion"
636 );
637 }
638
639 #[tokio::test]
644 async fn a_pre_upgrade_event_replays_literally() {
645 let agent = Arc::new(FnTextAgent::new("worker", |state| {
646 Ok(format!(
647 "{:?}/{:?}/{:?}",
648 state.get::<String>("adk:removed:literal"),
649 state.get::<Vec<String>>(EventActions::REMOVED_KEYS),
650 state.get::<String>("survivor"),
651 ))
652 }));
653 let runner = TextRunner::new(agent, "test-app");
654 let sessions = runner.session_service_ref();
655 let session = sessions.create_session("test-app", "user-1").await.unwrap();
656
657 let mut delta = std::collections::HashMap::new();
660 delta.insert(
661 "adk:removed:literal".to_string(),
662 serde_json::json!("laddered"),
663 );
664 delta.insert(
665 EventActions::REMOVED_KEYS.to_string(),
666 serde_json::json!(["survivor"]),
667 );
668 delta.insert("survivor".to_string(), serde_json::json!("still here"));
669 let legacy = Event::new("worker", Some("legacy".into()))
670 .with_actions(EventActions::state_delta(delta));
671 assert!(
672 !legacy.actions.is_format_marked(),
673 "the fixture must look like a pre-upgrade event"
674 );
675 sessions.append_event(&session.id, legacy).await.unwrap();
676
677 let peeked = runner
678 .run("peek", "user-1", Some(&session.id))
679 .await
680 .unwrap();
681 assert_eq!(
682 peeked, "Some(\"laddered\")/Some([\"survivor\"])/Some(\"still here\")",
683 "a pre-upgrade event was decoded under the new rules"
684 );
685 }
686
687 #[tokio::test]
688 async fn a_deliberately_stored_json_null_survives_replay() {
689 let agent = Arc::new(FnTextAgent::new("worker", |state| {
693 let input: String = state.get("input").unwrap_or_default();
694 if input == "store" {
695 let _ = state.set("maybe", serde_json::Value::Null);
696 }
697 Ok(format!("present: {}", state.contains("maybe")))
698 }));
699 let runner = TextRunner::new(agent, "test-app");
700
701 runner.run("store", "user-1", None).await.unwrap();
702 let sessions = runner
703 .session_service_ref()
704 .list_sessions("test-app", "user-1")
705 .await
706 .unwrap();
707 let sid = sessions[0].id.clone();
708
709 let peeked = runner.run("peek", "user-1", Some(&sid)).await.unwrap();
710 assert_eq!(
711 peeked, "present: true",
712 "a stored JSON null must survive replay, not be read as a deletion"
713 );
714 }
715
716 #[tokio::test]
717 async fn run_stream_drains_to_same_result_as_run() {
718 let runner = TextRunner::new(echo_agent(), "test-app");
719
720 let mut stream = runner.run_stream("Hi", "user-1", None);
722 let mut last = None;
723 while let Some(item) = stream.next().await {
724 if let RunEvent::Event(e) = item
725 && e.author == "echo"
726 {
727 last = e.content.clone();
728 }
729 }
730 assert_eq!(last.as_deref(), Some("Echo: Hi"));
731 }
732
733 #[tokio::test]
734 async fn run_stream_persists_events_like_run() {
735 let runner = TextRunner::new(echo_agent(), "test-app");
736 let mut stream = runner.run_stream("Hello", "user-1", None);
737 while stream.next().await.is_some() {}
738 drop(stream);
739
740 let sessions = runner
741 .session_service_ref()
742 .list_sessions("test-app", "user-1")
743 .await
744 .unwrap();
745 assert_eq!(sessions.len(), 1);
746 let events = runner
747 .session_service_ref()
748 .get_events(&sessions[0].id)
749 .await
750 .unwrap();
751 assert_eq!(events.len(), 2);
752 assert_eq!(events[0].author, "user");
753 assert_eq!(events[1].author, "echo");
754 }
755
756 #[tokio::test]
757 async fn run_stream_emits_error_for_missing_session() {
758 let runner = TextRunner::new(echo_agent(), "test-app");
759 let fake_id = SessionId::new();
760 let mut stream = runner.run_stream("Hello", "user-1", Some(&fake_id));
761
762 let mut saw_error = false;
763 while let Some(item) = stream.next().await {
764 if let RunEvent::Error(_) = item {
765 saw_error = true;
766 }
767 }
768 assert!(saw_error);
769 }
770}