1use std::sync::Arc;
17use std::sync::atomic::{AtomicBool, Ordering};
18
19use pi_agent::{AgentEvent, AgentEventSubscription};
20use tokio::task::JoinHandle;
21use tokio_util::sync::CancellationToken;
22
23use super::AgentSession;
24use super::events::AgentSessionEvent;
25
26pub(super) struct EventPump {
28 pub cancel: CancellationToken,
30 pub join: JoinHandle<()>,
32 pub active: Arc<AtomicBool>,
34}
35
36impl AgentSession {
37 pub(super) fn spawn_event_pump(self: &Arc<Self>) -> EventPump {
39 self.spawn_event_pump_with_subscription(self.agent.subscribe())
40 }
41
42 fn spawn_event_pump_with_subscription(
43 self: &Arc<Self>,
44 mut rx: AgentEventSubscription,
45 ) -> EventPump {
46 let cancel = CancellationToken::new();
47 let active = Arc::new(AtomicBool::new(true));
48 let session = Arc::clone(self);
49 let cancel_child = cancel.clone();
50 let active_flag = Arc::clone(&active);
51 let wait_cancel = self.lock_inner().agent_end_wait_cancel.clone();
52
53 let join = tokio::spawn(async move {
54 let mut lag_recorded = false;
55 loop {
56 tokio::select! {
57 () = cancel_child.cancelled() => break,
58 event = rx.recv() => {
59 let Some(event) = event else {
60 break;
61 };
62 if rx.is_lagged() && !lag_recorded {
63 lag_recorded = true;
70 session.record_session_error(
71 crate::core::sessions::SessionError::Io {
72 path: "agent event subscription".to_owned(),
73 source: std::io::Error::other(
74 "agent event subscription lagged; run lifecycle is incomplete",
75 ),
76 },
77 );
78 session.agent.abort();
79 }
80 session.process_agent_event(event).await;
81 }
82 }
83 }
84 wait_cancel.cancel();
85 active_flag.store(false, Ordering::SeqCst);
86 });
87
88 EventPump {
89 cancel,
90 join,
91 active,
92 }
93 }
94
95 pub(super) fn disconnect_from_agent(&self) {
97 self.lock_inner().agent_end_wait_cancel.cancel();
98 if let Some(pump) = self.take_pump() {
99 pump.cancel.cancel();
100 pump.join.abort();
103 }
104 }
105
106 pub(super) fn reconnect_to_agent(self: &Arc<Self>) {
108 if self.pump_is_active() {
109 return;
110 }
111 self.lock_inner().agent_end_wait_cancel = CancellationToken::new();
112 let pump = self.spawn_event_pump();
113 self.store_pump(pump);
114 }
115
116 async fn process_agent_event(self: &Arc<Self>, event: AgentEvent) {
118 let is_agent_end = matches!(&event, AgentEvent::AgentEnd { .. });
119 if matches!(&event, AgentEvent::MessageStart { message } if message.role() == "user")
120 && let Err(error) = self
121 .handle_agent_event_side_effects(&event, &AgentSessionEvent::AgentStart)
122 .await
123 {
124 self.record_session_error(error);
128 self.agent.abort();
129 return;
130 }
131
132 let (public, persistence_event) = match event {
133 AgentEvent::MessageUpdate {
134 message,
135 assistant_message_event,
136 } => {
137 let runner = self.hooks.runner();
138 if runner.has_handlers("message_update")
139 && let Err(error) = runner
140 .emit_message_update_delta(assistant_message_event.as_ref())
141 .await
142 {
143 runner.emit_error(error.to_string());
144 }
145 (
146 AgentSessionEvent::MessageUpdate {
147 message,
148 assistant_message_event,
149 },
150 None,
151 )
152 }
153 AgentEvent::MessageEnd { message } => {
154 let runner = self.hooks.runner();
155 let replacement = if runner.has_handlers("message_end") {
156 match runner.emit_message_end(message.clone()).await {
157 Ok(replacement) => replacement,
158 Err(error) => {
159 runner.emit_error(error.to_string());
160 None
161 }
162 }
163 } else {
164 None
165 };
166 let public_message = replacement.map_or_else(
167 || message.clone(),
168 |replacement| {
169 self.apply_message_end_replacement(message.clone(), Some(replacement))
170 },
171 );
172 (
173 AgentSessionEvent::MessageEnd {
174 message: public_message,
175 },
176 Some(AgentEvent::MessageEnd { message }),
177 )
178 }
179 event => {
180 let public = self.map_agent_event_for_public(event);
181 let runner = self.hooks.runner();
182 if runner.has_handlers(public.type_name())
183 && let Err(error) = runner.emit(public.clone()).await
184 {
185 runner.emit_error(error.to_string());
186 }
187 (public, None)
188 }
189 };
190
191 self.emit_public_awaited(&public).await;
192
193 if let Some(event) = persistence_event
194 && let Err(error) = self.handle_agent_event_side_effects(&event, &public).await
195 {
196 self.record_session_error(error);
199 self.agent.abort();
200 }
201
202 if is_agent_end {
203 let notify = {
204 let mut inner = self.lock_inner();
205 inner.processed_agent_ends = inner.processed_agent_ends.saturating_add(1);
206 Arc::clone(&inner.agent_end_notify)
207 };
208 notify.notify_waiters();
209 }
210 }
211
212 pub(super) fn processed_agent_end_count(&self) -> u64 {
213 self.lock_inner().processed_agent_ends
214 }
215
216 pub(super) async fn wait_for_processed_agent_end(&self, before: u64) -> bool {
217 loop {
218 let (notified, cancelled) = {
219 let inner = self.lock_inner();
220 if inner.processed_agent_ends > before {
221 return true;
222 }
223 (
224 Arc::clone(&inner.agent_end_notify).notified_owned(),
225 inner.agent_end_wait_cancel.clone(),
226 )
227 };
228 tokio::select! {
229 () = notified => {}
230 () = cancelled.cancelled() => return false,
231 }
232 }
233 }
234
235 fn map_agent_event_for_public(&self, event: AgentEvent) -> AgentSessionEvent {
237 let will_retry = match &event {
238 AgentEvent::AgentEnd { messages } => self.will_retry_after_agent_end(messages),
239 _ => false,
240 };
241 AgentSessionEvent::from_agent_event(event, will_retry)
242 }
243
244 pub(super) fn will_retry_after_agent_end(&self, messages: &[pi_agent::AgentMessage]) -> bool {
249 let inner = self.lock_inner();
250 if !inner.auto_retry_enabled || inner.retry_attempt >= inner.max_retries {
251 return false;
252 }
253 for message in messages.iter().rev() {
254 if message.role() == "assistant" {
255 if let Some(pi_ai::Message::Assistant(assistant)) = message.as_llm() {
256 return Self::is_retryable_error(assistant);
257 }
258 return false;
259 }
260 }
261 false
262 }
263
264 pub async fn emit_agent_settled(self: &Arc<Self>) {
269 {
270 let mut inner = self.lock_inner();
271 if !inner.is_agent_run_active {
272 return;
274 }
275 inner.is_agent_run_active = false;
276 }
277
278 let runner = self.hooks.runner();
279 let _ = runner.emit(AgentSessionEvent::AgentSettled).await;
280 self.emit_public_awaited(&AgentSessionEvent::AgentSettled)
281 .await;
282 self.resolve_idle_waiters();
283 }
284
285 #[cfg(test)]
286 pub(super) fn mark_agent_run_active(&self) {
288 let mut inner = self.lock_inner();
289 inner.is_agent_run_active = true;
290 }
291
292 pub(super) fn resolve_idle_waiters(&self) {
294 let inner = self.lock_inner();
295 if inner.is_agent_run_active {
296 return;
297 }
298 inner.idle_notify.notify_waiters();
299 }
300}
301
302#[cfg(test)]
303mod tests {
304 use super::*;
305 use futures::stream::{self, BoxStream, StreamExt};
306 use pi_agent::{AgentEventSink, AgentState, EventSink};
307 use pi_ai::{
308 AssistantMessageEvent, Context, Model, ModelCost, ModelInput, Provider, ProviderError,
309 StreamOptions,
310 };
311
312 type TestResult = Result<(), Box<dyn std::error::Error>>;
313
314 #[derive(Clone)]
315 struct StubProvider;
316
317 impl Provider for StubProvider {
318 fn stream(
319 &self,
320 _model: &Model,
321 _context: Context,
322 _options: StreamOptions,
323 ) -> BoxStream<'static, Result<AssistantMessageEvent, ProviderError>> {
324 stream::empty().boxed()
325 }
326 }
327
328 fn model() -> Model {
329 Model {
330 id: "model".into(),
331 name: "model".into(),
332 api: "test".into(),
333 provider: "test".into(),
334 base_url: String::new(),
335 reasoning: false,
336 thinking_level_map: None,
337 input: vec![ModelInput::Text],
338 cost: ModelCost::default(),
339 context_window: 8_192,
340 max_tokens: 1_024,
341 headers: None,
342 compat: None,
343 extra: std::collections::BTreeMap::new(),
344 }
345 }
346
347 fn assistant_message() -> pi_agent::AgentMessage {
348 let mut assistant = pi_ai::AssistantMessage::new("test", "test", "model", 0);
349 assistant.stop_reason = pi_ai::StopReason::Stop;
350 pi_agent::AgentMessage::Llm(Box::new(pi_ai::Message::Assistant(assistant)))
351 }
352
353 fn update_event(index: i64) -> AgentEvent {
354 let partial = pi_ai::AssistantMessage::new("test", "test", "model", index);
355 AgentEvent::MessageUpdate {
356 message: assistant_message(),
357 assistant_message_event: Box::new(AssistantMessageEvent::Start { partial }),
358 }
359 }
360
361 #[tokio::test]
362 async fn lagged_subscription_drains_retained_terminals_before_settle() -> TestResult {
363 let config =
364 super::super::AgentSessionConfig::test_config(Arc::new(StubProvider), model())?;
365 let session = super::super::AgentSession::new(config)?;
366 let observed = Arc::new(std::sync::Mutex::new(Vec::new()));
367 let observed_clone = Arc::clone(&observed);
368 let _unsub = session.subscribe(move |event| {
369 observed_clone
370 .lock()
371 .unwrap_or_else(std::sync::PoisonError::into_inner)
372 .push(event.type_name().to_owned());
373 });
374
375 session.mark_agent_run_active();
377 let sink = AgentEventSink::new(Arc::new(std::sync::Mutex::new(AgentState::new())));
378 let rx = sink.subscribe_with_capacity(2);
379 let ends_before = session.processed_agent_end_count();
380 for index in 0..4 {
383 sink.emit(update_event(index));
384 }
385 sink.emit(AgentEvent::MessageEnd {
386 message: assistant_message(),
387 });
388 sink.emit(AgentEvent::AgentEnd {
389 messages: Vec::new(),
390 });
391 drop(sink);
392
393 let pump = session.spawn_event_pump_with_subscription(rx);
394 pump.join.await?;
395
396 let error = session
397 .take_session_error()
398 .ok_or("lag must record a typed session error")?;
399 assert!(error.to_string().contains("subscription lagged"), "{error}");
400 assert!(session.take_session_error().is_none(), "error is one-shot");
401 assert_eq!(
402 session.processed_agent_end_count(),
403 ends_before + 1,
404 "retained agent_end must still be processed after lag"
405 );
406
407 let snapshot = observed
408 .lock()
409 .unwrap_or_else(std::sync::PoisonError::into_inner)
410 .clone();
411 assert!(
412 snapshot.contains(&"message_end".to_owned()),
413 "retained message_end must be published: {snapshot:?}"
414 );
415 assert!(
416 snapshot.contains(&"agent_end".to_owned()),
417 "retained agent_end must be published: {snapshot:?}"
418 );
419 assert!(
420 !snapshot.contains(&"agent_settled".to_owned()),
421 "the pump must never settle on lag: {snapshot:?}"
422 );
423
424 session.emit_agent_settled().await;
427 let snapshot = observed
428 .lock()
429 .unwrap_or_else(std::sync::PoisonError::into_inner)
430 .clone();
431 let end = snapshot
432 .iter()
433 .position(|name| name == "agent_end")
434 .ok_or("agent_end position")?;
435 let settled = snapshot
436 .iter()
437 .position(|name| name == "agent_settled")
438 .ok_or("agent_settled position")?;
439 assert!(end < settled, "AgentEnd must precede settle: {snapshot:?}");
440 assert_eq!(
441 snapshot
442 .iter()
443 .filter(|name| *name == "agent_settled")
444 .count(),
445 1,
446 "exactly one settle: {snapshot:?}"
447 );
448 Ok(())
449 }
450}