Skip to main content

pi/core/agent_session/
subscribe.rs

1//! Bounded agent-event pump.
2//!
3//! Spawns exactly one task that drains `Agent::subscribe()`. Any observable
4//! subscription lag becomes a typed session failure and aborts/settles the run
5//! rather than persisting a partial lifecycle. For each retained event it:
6//! 1. runs pre-public side effects (`message_start` queue dequeue)
7//! 2. awaits the extension handler (compact deltas for streaming updates)
8//! 3. applies `message_end` replacement to agent state when present
9//! 4. emits the public `AgentSessionEvent` (listeners called without holding locks)
10//! 5. persists `message_end` / emits auto-retry success
11//!
12//! Session-level `agent_settled` is emitted exactly once when the session run
13//! flag flips false (after retries / follow-ups / auto-compaction). The same
14//! guarded entry point also unblocks waiters after persistence or pump failure.
15
16use 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
26/// Handle for the background event pump.
27pub(super) struct EventPump {
28    /// Cancellation token cancelled on dispose / disconnect.
29    pub cancel: CancellationToken,
30    /// Join handle for the pump task.
31    pub join: JoinHandle<()>,
32    /// True while the pump is the active agent subscription.
33    pub active: Arc<AtomicBool>,
34}
35
36impl AgentSession {
37    /// Spawn the single bounded event pump for this session.
38    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                            // Record once and abort, but keep draining: the
64                            // bounded subscription retains the newest
65                            // message_end/agent_end, and only a processed
66                            // agent_end may drive the single settle. If the
67                            // subscription closes instead, the loop breaks and
68                            // wait_cancel unblocks the prompt barrier.
69                            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    /// Disconnect the pump without disposing the session (compaction pause).
96    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            // Detach; do not await join here (may be called from async context
101            // that must not block). The task exits promptly on cancel.
102            pump.join.abort();
103        }
104    }
105
106    /// Reconnect a pump after disconnect.
107    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    /// Process one agent event through extension → public → persistence.
117    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            // Record the typed failure and abort the run, but never settle
125            // here: the aborted run still emits `agent_end`, and the prompt
126            // lifecycle settles exactly once after processing it.
127            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            // Same invariant as the message_start path: record + abort, and
197            // let the eventual `agent_end` drive the single settle.
198            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    /// Map a core agent event into a public session event.
236    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    /// Whether auto-retry will continue after this `agent_end`.
245    ///
246    /// Uses the same retry.rs classifier and runtime attempt/settings gates as
247    /// actual retry execution so public prediction cannot drift.
248    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    /// Emit `agent_settled` exactly once for the current session-level run.
265    ///
266    /// Callers (prompt lifecycle) invoke this after retries/follow-ups complete.
267    /// Extension handler is awaited before public listeners.
268    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                // Already settled; do not double-emit.
273                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    /// Mark the session-level run active (prompt start).
287    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    /// Resolve waiters blocked in `wait_for_idle` when session is idle.
293    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        // Mirror the prompt lifecycle: the run flag gates the single settle.
376        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        // Overflow with streaming updates; the bounded subscription must keep
381        // the newest message_end and agent_end deliverable.
382        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        // Settle exactly once, strictly after the processed agent_end —
425        // mirroring the prompt lifecycle that owns the settle.
426        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}