Skip to main content

pi/core/agent_session/
prompt.rs

1//! Prompt lifecycle: preflight, steer/follow-up queues, run loop, and post-run
2//! continuation (retry → compaction → queued messages).
3//!
4//! [`AgentSession::prompt`] runs the preflight then calls
5//! [`AgentSession::run_agent_prompt`] which loops: `agent.prompt(messages)`,
6//! then `while handle_post_agent_run` (retry → compaction → queued messages)
7//! triggers another `agent.continue_run()`. `emit_agent_settled` fires
8//! exactly once in the finally block.
9
10use std::sync::Arc;
11
12use pi_agent::{AgentMessage, user_text};
13use pi_ai::{AssistantMessage, ImageContent};
14
15use super::events::AgentSessionEvent;
16use super::{AgentSession, BeforeAgentStartResult};
17use crate::core::agent_session_services::{
18    format_no_api_key_found_message, format_no_model_selected_message,
19    format_oauth_auth_failed_message,
20};
21use crate::core::messages::CustomMessageContent;
22use crate::core::model_runtime::ModelRuntime;
23use crate::core::resources::frontmatter::strip_frontmatter;
24use crate::core::resources::prompts::expand_prompt_template;
25
26/// Upper bound on waiting for a failed run's already-queued `agent_end`.
27///
28/// Ordinary failed runs never wait this long: their `agent_end` is queued
29/// before the run call returns, so the barrier resolves on the next pump
30/// iteration. This bound only fires for a future pre-start rejection that
31/// [`run_emits_agent_end`] does not classify, keeping worst-case added
32/// latency well under the interactive p95 budget.
33const FAILED_RUN_END_BARRIER_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(1);
34
35/// How a streaming prompt is queued.
36#[derive(Clone, Copy, Debug, Eq, PartialEq)]
37pub enum StreamingBehavior {
38    /// Inject before the next LLM call (steering queue).
39    Steer,
40    /// Inject after the current run finishes (follow-up queue).
41    FollowUp,
42}
43
44impl StreamingBehavior {
45    /// Wire string matching TypeScript (`"steer"` / `"followUp"`).
46    #[must_use]
47    pub fn as_str(&self) -> &'static str {
48        match self {
49            Self::Steer => "steer",
50            Self::FollowUp => "followUp",
51        }
52    }
53}
54
55/// Callback invoked once during `prompt()` preflight to signal accept/reject.
56///
57/// `true` is called when the prompt will be run (or queued), `false` on a
58/// synchronous error. RPC consumers use this to emit exactly one response
59/// per `prompt()` call.
60pub type PreflightCallback = Arc<dyn Fn(bool) + Send + Sync>;
61
62/// Configuration for one [`AgentSession::prompt`] invocation.
63#[derive(Clone)]
64pub struct PromptOptions {
65    /// Image attachments for the user message.
66    pub images: Vec<ImageContent>,
67    /// When the session is already streaming, routes the message to the
68    /// steering or follow-up queue. Required while streaming.
69    pub streaming_behavior: Option<StreamingBehavior>,
70    /// Origin label forwarded to extension `input` events.
71    pub source: Option<String>,
72    /// Expand extension commands, `/skill:name`, and prompt templates.
73    pub expand_prompt_templates: bool,
74    /// Optional single-shot accept signal (see [`PreflightCallback`]).
75    pub preflight_result: Option<PreflightCallback>,
76}
77
78impl Default for PromptOptions {
79    fn default() -> Self {
80        Self {
81            images: Vec::new(),
82            streaming_behavior: None,
83            source: None,
84            expand_prompt_templates: true,
85            preflight_result: None,
86        }
87    }
88}
89
90impl PromptOptions {
91    /// Build defaults.
92    #[must_use]
93    pub fn new() -> Self {
94        Self::default()
95    }
96}
97
98/// Delivery mode for [`AgentSession::send_custom_message`].
99#[derive(Clone, Copy, Debug, Eq, PartialEq)]
100pub enum DeliverAs {
101    /// Steering queue (injected before the next LLM call).
102    Steer,
103    /// Follow-up queue (injected after the current run).
104    FollowUp,
105    /// Buffered; appended alongside the next user prompt.
106    NextTurn,
107}
108
109/// Input for [`AgentSession::send_custom_message`].
110#[derive(Clone, Debug)]
111pub struct CustomMessageInput {
112    /// Extension-defined discriminant.
113    pub custom_type: String,
114    /// User-visible / LLM content.
115    pub content: CustomMessageContent,
116    /// Whether the interactive UI should render this message.
117    pub display: bool,
118    /// Opaque details preserved on the transcript entry.
119    pub details: Option<serde_json::Value>,
120}
121
122/// Errors produced by the prompt lifecycle.
123#[derive(Debug, thiserror::Error)]
124pub enum PromptError {
125    /// Human-readable error (no model / no auth / queue / extension command).
126    #[error("{0}")]
127    Message(String),
128    /// Underlying agent run failure.
129    #[error(transparent)]
130    Agent(#[from] pi_agent::AgentLoopError),
131    /// Session persistence failure while preparing or settling the prompt.
132    #[error(transparent)]
133    Session(#[from] crate::core::sessions::SessionError),
134}
135
136impl PromptError {
137    #[must_use]
138    fn msg(s: impl Into<String>) -> Self {
139        Self::Message(s.into())
140    }
141}
142
143fn map_bash_flush_error(err: super::bash::BashExecError) -> PromptError {
144    match err {
145        super::bash::BashExecError::Session(err) => PromptError::Session(err),
146        super::bash::BashExecError::Execution { message, .. } => PromptError::Message(message),
147    }
148}
149
150struct RunAdmission {
151    session: Arc<AgentSession>,
152    armed: bool,
153}
154
155impl RunAdmission {
156    fn disarm(&mut self) {
157        self.armed = false;
158    }
159}
160
161impl Drop for RunAdmission {
162    fn drop(&mut self) {
163        if !self.armed {
164            return;
165        }
166        {
167            let mut inner = self.session.lock_inner();
168            inner.is_agent_run_active = false;
169        }
170        self.session.resolve_idle_waiters();
171    }
172}
173
174enum PreflightOutcome {
175    Run {
176        messages: Vec<AgentMessage>,
177        admission: RunAdmission,
178    },
179    Handled,
180    Queued,
181}
182
183impl AgentSession {
184    /// Send a prompt to the agent. Handles extension commands, input transforms,
185    /// skill/template expansion, streaming-queue routing, model + auth validation,
186    /// pre-prompt compaction, `before_agent_start` injection, and the run loop.
187    ///
188    /// # Errors
189    ///
190    /// - [`PromptError::Message`] for no-model, no-auth, concurrent-streaming
191    ///   guard, or queue slash-command rejection.
192    /// - [`PromptError::Agent`] when the underlying agent run fails.
193    /// - [`PromptError::Session`] when transcript persistence fails.
194    pub async fn prompt(
195        self: &Arc<Self>,
196        text: &str,
197        options: PromptOptions,
198    ) -> Result<(), PromptError> {
199        self.prompt_inner(text, options).await
200    }
201
202    async fn prompt_inner(
203        self: &Arc<Self>,
204        text: &str,
205        mut options: PromptOptions,
206    ) -> Result<(), PromptError> {
207        let expand = options.expand_prompt_templates;
208        let preflight = options.preflight_result.take();
209        let call_preflight = |ok: bool| {
210            if let Some(cb) = &preflight {
211                cb(ok);
212            }
213        };
214
215        match self.prompt_preflight(text, &mut options, expand).await {
216            Ok(PreflightOutcome::Run {
217                messages,
218                admission,
219            }) => {
220                call_preflight(true);
221                self.run_agent_prompt_admitted(messages, admission).await?;
222                Ok(())
223            }
224            Ok(PreflightOutcome::Handled | PreflightOutcome::Queued) => {
225                call_preflight(true);
226                Ok(())
227            }
228            Err(err) => {
229                call_preflight(false);
230                Err(err)
231            }
232        }
233    }
234
235    /// Queue a steering message; errors when `text` is a registered extension
236    /// command (commands cannot be queued).
237    ///
238    /// # Errors
239    ///
240    /// [`PromptError::Message`] for extension-command rejection.
241    pub fn steer(&self, text: &str, images: Vec<ImageContent>) -> Result<(), PromptError> {
242        self.check_not_extension_command(text)?;
243        let expanded = self.expand_text(text);
244        self.queue_steer(&expanded, images);
245        Ok(())
246    }
247
248    /// Queue a follow-up message; errors when `text` is a registered extension
249    /// command.
250    ///
251    /// # Errors
252    ///
253    /// [`PromptError::Message`] for extension-command rejection.
254    pub fn follow_up(&self, text: &str, images: Vec<ImageContent>) -> Result<(), PromptError> {
255        self.check_not_extension_command(text)?;
256        let expanded = self.expand_text(text);
257        self.queue_follow_up(&expanded, images);
258        Ok(())
259    }
260
261    /// Send a custom message (extension-injected transcript entry). Delivery is
262    /// selected by `deliver_as` and current streaming state.
263    ///
264    /// # Errors
265    ///
266    /// Returns [`PromptError`] when starting a requested agent turn fails or
267    /// when the idle-path durable session append fails (no live state or
268    /// public event is published in that case).
269    pub async fn send_custom_message(
270        self: &Arc<Self>,
271        message: CustomMessageInput,
272        trigger_turn: bool,
273        deliver_as: Option<DeliverAs>,
274    ) -> Result<(), PromptError> {
275        let app_message = build_custom_agent_message(&message);
276
277        match deliver_as {
278            Some(DeliverAs::NextTurn) => {
279                self.lock_inner()
280                    .pending_next_turn_messages
281                    .push(app_message);
282            }
283            _ if self.is_session_streaming() => match deliver_as {
284                Some(DeliverAs::FollowUp) => self.agent.follow_up(app_message),
285                _ => self.agent.steer(app_message),
286            },
287            _ if trigger_turn => {
288                self.run_agent_prompt(vec![app_message]).await?;
289            }
290            _ => {
291                // Durable append first: live agent state and public events are
292                // published only for entries the session file actually holds.
293                {
294                    let mut sm = self.session_manager.lock().await;
295                    sm.append_custom_message_entry(
296                        &message.custom_type,
297                        &message.content,
298                        message.display,
299                        message.details.clone(),
300                    )
301                    .map_err(PromptError::Session)?;
302                }
303                self.agent.push_message(app_message.clone());
304                self.emit_public(AgentSessionEvent::MessageStart {
305                    message: app_message.clone(),
306                });
307                self.emit_public(AgentSessionEvent::MessageEnd {
308                    message: app_message,
309                });
310            }
311        }
312        Ok(())
313    }
314
315    /// Send a user message. While idle this triggers a new turn; while
316    /// streaming it queues per `deliver_as`.
317    ///
318    /// # Errors
319    ///
320    /// Returns [`PromptError`] when prompt validation, extension handling, or
321    /// the underlying agent run fails.
322    pub async fn send_user_message(
323        self: &Arc<Self>,
324        text: &str,
325        images: Vec<ImageContent>,
326        deliver_as: Option<DeliverAs>,
327    ) -> Result<(), PromptError> {
328        let streaming_behavior = deliver_as.map(|d| match d {
329            DeliverAs::Steer | DeliverAs::NextTurn => StreamingBehavior::Steer,
330            DeliverAs::FollowUp => StreamingBehavior::FollowUp,
331        });
332        self.prompt(
333            text,
334            PromptOptions {
335                images,
336                streaming_behavior,
337                source: Some("extension".to_owned()),
338                expand_prompt_templates: false,
339                preflight_result: None,
340            },
341        )
342        .await
343    }
344
345    // -----------------------------------------------------------------
346    // Preflight
347    // -----------------------------------------------------------------
348
349    async fn prompt_preflight(
350        self: &Arc<Self>,
351        text: &str,
352        options: &mut PromptOptions,
353        expand: bool,
354    ) -> Result<PreflightOutcome, PromptError> {
355        // 1. Extension command dispatch.
356        if expand && text.starts_with('/') && self.try_execute_extension_command(text).await? {
357            return Ok(PreflightOutcome::Handled);
358        }
359
360        // 2. Input event transform.
361        let current_images = std::mem::take(&mut options.images);
362        let (current_text, current_images, handled) = self
363            .transform_input(text.to_owned(), current_images, options)
364            .await?;
365        if handled {
366            options.images = current_images;
367            return Ok(PreflightOutcome::Handled);
368        }
369
370        // 3. Skill / template expansion.
371        let expanded_text = if expand {
372            self.expand_text(&current_text)
373        } else {
374            current_text
375        };
376
377        // 4. Streaming: queue via steer/followUp.
378        let Some(admission) = self.reserve_run_admission() else {
379            let behavior = options.streaming_behavior.ok_or_else(|| {
380                PromptError::msg(
381                    "Agent is already processing. Specify streamingBehavior \
382                     ('steer' or 'followUp') to queue the message.",
383                )
384            })?;
385            match behavior {
386                StreamingBehavior::FollowUp => {
387                    self.queue_follow_up(&expanded_text, current_images);
388                }
389                StreamingBehavior::Steer => {
390                    self.queue_steer(&expanded_text, current_images);
391                }
392            }
393            return Ok(PreflightOutcome::Queued);
394        };
395
396        // 5. Flush any pending bash messages before validation.
397        self.flush_pending_bash_messages()
398            .await
399            .map_err(map_bash_flush_error)?;
400
401        // 6. Validate model.
402        let model = self.model();
403        if is_no_model(&model) {
404            return Err(PromptError::Message(format_no_model_selected_message()));
405        }
406
407        // 7. Validate auth.
408        if let Some(runtime) = self.try_model_runtime() {
409            let provider = &model.provider;
410            let has_auth = runtime.has_configured_auth(provider)
411                || runtime.check_auth(provider).await.is_some();
412            if !has_auth {
413                if runtime.is_using_oauth(provider) {
414                    return Err(PromptError::Message(format_oauth_auth_failed_message(
415                        provider,
416                    )));
417                }
418                return Err(PromptError::Message(format_no_api_key_found_message(
419                    provider,
420                )));
421            }
422        }
423
424        // 8. Pre-prompt compaction check (no agent.continue — sibling compaction
425        //    slice owns this; here we trigger only the pre-prompt pass).
426        if let Some(last_msg) = self.agent.last_assistant() {
427            self.check_compaction(&last_msg, false).await;
428        }
429
430        // 9. Build messages: user message + pending nextTurn.
431        let mut messages = Vec::new();
432        messages.push(user_text(&expanded_text, current_images.iter().cloned()));
433        {
434            let mut inner = self.lock_inner();
435            for msg in inner.pending_next_turn_messages.drain(..) {
436                messages.push(msg);
437            }
438        }
439
440        // 10. before_agent_start extension event.
441        let runner = self.hooks.runner();
442        let images_for_ext = if current_images.is_empty() {
443            None
444        } else {
445            serde_json::to_value(&current_images).ok()
446        };
447        let result = runner
448            .emit_before_agent_start(&expanded_text, images_for_ext)
449            .await
450            .map_err(|e| PromptError::msg(e.to_string()))?;
451        drop(runner);
452
453        self.apply_before_agent_start(result);
454
455        Ok(PreflightOutcome::Run {
456            messages,
457            admission,
458        })
459    }
460
461    async fn transform_input(
462        &self,
463        mut text: String,
464        mut images: Vec<ImageContent>,
465        options: &PromptOptions,
466    ) -> Result<(String, Vec<ImageContent>, bool), PromptError> {
467        let runner = self.hooks.runner();
468        if !runner.has_handlers("input") {
469            return Ok((text, images, false));
470        }
471
472        let streaming = if self.is_session_streaming() {
473            options.streaming_behavior.map(|behavior| behavior.as_str())
474        } else {
475            None
476        };
477        let source = options.source.as_deref().unwrap_or("interactive");
478        let images_value = if images.is_empty() {
479            None
480        } else {
481            serde_json::to_value(&images).ok()
482        };
483        let result = runner
484            .emit_input(&text, images_value, source, streaming)
485            .await
486            .map_err(|error| PromptError::msg(error.to_string()))?;
487
488        if result.handled {
489            return Ok((text, images, true));
490        }
491
492        if let Some(transformed_text) = result.text {
493            text = transformed_text;
494        }
495        if let Some(transformed_images) = result.images
496            && let Some(parsed) = parse_images_value(&transformed_images)
497        {
498            images = parsed;
499        }
500
501        Ok((text, images, false))
502    }
503
504    fn apply_before_agent_start(&self, result: Option<BeforeAgentStartResult>) {
505        let Some(result) = result else {
506            self.hooks.set_system_prompt_override(None);
507            let base = self.lock_inner().base_system_prompt.clone();
508            self.agent.set_system_prompt(base);
509            return;
510        };
511
512        if !result.messages.is_empty() {
513            let mut inner = self.lock_inner();
514            inner
515                .pending_next_turn_messages
516                .splice(0..0, result.messages);
517        }
518
519        if let Some(system_prompt) = result.system_prompt {
520            self.hooks
521                .set_system_prompt_override(Some(system_prompt.clone()));
522            self.agent.set_system_prompt(system_prompt);
523        } else {
524            self.hooks.set_system_prompt_override(None);
525            let base = self.lock_inner().base_system_prompt.clone();
526            self.agent.set_system_prompt(base);
527        }
528    }
529
530    // -----------------------------------------------------------------
531    // Run loop
532    // -----------------------------------------------------------------
533
534    async fn run_agent_prompt(
535        self: &Arc<Self>,
536        messages: Vec<AgentMessage>,
537    ) -> Result<(), PromptError> {
538        let admission = self.reserve_run_admission().ok_or_else(|| {
539            PromptError::msg(
540                "Agent is already processing. Specify streamingBehavior \
541                 ('steer' or 'followUp') to queue the message.",
542            )
543        })?;
544        self.run_agent_prompt_admitted(messages, admission).await
545    }
546
547    async fn run_agent_prompt_admitted(
548        self: &Arc<Self>,
549        messages: Vec<AgentMessage>,
550        mut admission: RunAdmission,
551    ) -> Result<(), PromptError> {
552        let result = self.run_agent_prompt_inner(messages).await;
553        // Flush any bash messages that arrived during the run before settle.
554        let flush_result = self
555            .flush_pending_bash_messages()
556            .await
557            .map_err(map_bash_flush_error);
558        self.hooks.set_system_prompt_override(None);
559        admission.disarm();
560        self.emit_agent_settled().await;
561        let pending_session_error = self.take_session_error();
562        if let Some(error) = pending_session_error {
563            return Err(PromptError::Session(error));
564        }
565        flush_result?;
566        result
567    }
568
569    async fn run_agent_prompt_inner(
570        self: &Arc<Self>,
571        messages: Vec<AgentMessage>,
572    ) -> Result<(), PromptError> {
573        let mut messages = messages;
574        {
575            let mut inner = self.lock_inner();
576            if !inner.pending_next_turn_messages.is_empty() {
577                let pending: Vec<_> = inner.pending_next_turn_messages.drain(..).collect();
578                messages.extend(pending);
579            }
580        }
581
582        // Capture assistant count BEFORE the run so the first response is seen
583        // as new and triggers retry/compaction/queued-message continuation.
584        // After prepare_retry / overflow compaction pop the trailing assistant,
585        // the count drops; re-baseline before continue_run so the next terminal
586        // assistant is observed (TS tracks this via `_lastAssistantMessage`).
587        let mut processed_count = self.assistant_count();
588        let mut processed_agent_ends = self.processed_agent_end_count();
589        let run = self.agent.prompt(messages).await;
590        self.observe_run_agent_end(&run, processed_agent_ends)
591            .await?;
592        if let Some(error) = self.take_session_error() {
593            return Err(PromptError::Session(error));
594        }
595        run?;
596        processed_agent_ends = self.processed_agent_end_count();
597        loop {
598            let current_count = self.assistant_count();
599            let new_assistant = if current_count > processed_count {
600                self.agent.last_assistant()
601            } else {
602                None
603            };
604            if !self.handle_post_agent_run(new_assistant).await? {
605                break;
606            }
607            // Re-baseline after pops from prepare_retry / overflow compaction.
608            processed_count = self.assistant_count();
609            let run = self.agent.continue_run().await;
610            self.observe_run_agent_end(&run, processed_agent_ends)
611                .await?;
612            if let Some(error) = self.take_session_error() {
613                return Err(PromptError::Session(error));
614            }
615            run?;
616            processed_agent_ends = self.processed_agent_end_count();
617        }
618        Ok(())
619    }
620
621    /// Await the processed `agent_end` barrier for one run outcome.
622    ///
623    /// Successful runs require the barrier: a pump disconnect before the end
624    /// event is a hard prompt error. Failed runs that started a lifecycle have
625    /// already queued their synthesized `agent_end` before returning, so the
626    /// barrier resolves immediately; the bounded timeout guards against a
627    /// pre-start rejection this classifier does not know about, so a
628    /// misclassification can never hang the prompt lifecycle.
629    async fn observe_run_agent_end(
630        &self,
631        run: &Result<(), pi_agent::AgentLoopError>,
632        processed_agent_ends: u64,
633    ) -> Result<(), PromptError> {
634        if run.is_ok() {
635            if !self
636                .wait_for_processed_agent_end(processed_agent_ends)
637                .await
638            {
639                return Err(PromptError::msg(
640                    "Agent event pump disconnected before agent_end",
641                ));
642            }
643            return Ok(());
644        }
645        if run_emits_agent_end(run) {
646            let _ = tokio::time::timeout(
647                FAILED_RUN_END_BARRIER_TIMEOUT,
648                self.wait_for_processed_agent_end(processed_agent_ends),
649            )
650            .await;
651        }
652        Ok(())
653    }
654
655    /// Post-run check: retry → compaction → queued messages.
656    ///
657    /// Returns `true` when the caller should `continue_run`.
658    async fn handle_post_agent_run(
659        self: &Arc<Self>,
660        msg: Option<AssistantMessage>,
661    ) -> Result<bool, PromptError> {
662        let Some(msg) = msg else {
663            return Ok(false);
664        };
665
666        if Self::is_retryable_error(&msg) && self.prepare_retry(&msg).await {
667            return Ok(true);
668        }
669
670        self.emit_retry_exhausted(&msg);
671
672        // Compaction check after retry handling. skip_aborted_check is `true`
673        // here because we only run after a real agent message_end, not an
674        // aborted user prompt.
675        if self.check_compaction(&msg, true).await {
676            return Ok(true);
677        }
678
679        Ok(self.agent.has_queued_messages())
680    }
681
682    // -----------------------------------------------------------------
683    // Queue helpers
684    // -----------------------------------------------------------------
685
686    fn queue_steer(&self, text: &str, images: Vec<ImageContent>) {
687        self.mirror_steering_push(text.to_owned());
688        self.agent.steer(user_text(text, images));
689    }
690
691    fn queue_follow_up(&self, text: &str, images: Vec<ImageContent>) {
692        self.mirror_follow_up_push(text.to_owned());
693        self.agent.follow_up(user_text(text, images));
694    }
695
696    // -----------------------------------------------------------------
697    // Extension commands + skill/template expansion
698    // -----------------------------------------------------------------
699
700    async fn try_execute_extension_command(&self, text: &str) -> Result<bool, PromptError> {
701        let Some((name, args)) = parse_slash_command(text) else {
702            return Ok(false);
703        };
704        let runner = self.hooks.runner();
705        if !runner.has_command(name) {
706            return Ok(false);
707        }
708        match runner.execute_command(name, args).await {
709            Ok(handled) => {
710                if !handled {
711                    return Ok(false);
712                }
713            }
714            Err(err) => {
715                runner.emit_error(format!("command:{name}: {err}"));
716            }
717        }
718        Ok(true)
719    }
720
721    fn check_not_extension_command(&self, text: &str) -> Result<(), PromptError> {
722        if !text.starts_with('/') {
723            return Ok(());
724        }
725        if let Some((name, _)) = parse_slash_command(text) {
726            let runner = self.hooks.runner();
727            if runner.has_command(name) {
728                return Err(PromptError::msg(format!(
729                    "Extension command \"/{name}\" cannot be queued. Use prompt() \
730                     or execute the command when not streaming."
731                )));
732            }
733        }
734        Ok(())
735    }
736
737    fn expand_text(&self, text: &str) -> String {
738        let expanded = self.expand_skill_command(text);
739        let templates = self
740            .prompt_templates
741            .lock()
742            .unwrap_or_else(std::sync::PoisonError::into_inner)
743            .clone();
744        expand_prompt_template(&expanded, &templates)
745    }
746
747    fn expand_skill_command(&self, text: &str) -> String {
748        if !text.starts_with("/skill:") {
749            return text.to_owned();
750        }
751        let rest = &text["/skill:".len()..];
752        let (skill_name, args) = match rest.find(' ') {
753            Some(idx) => (&rest[..idx], rest[idx + 1..].trim()),
754            None => (rest, ""),
755        };
756
757        let skills = self
758            .skills
759            .lock()
760            .unwrap_or_else(std::sync::PoisonError::into_inner)
761            .clone();
762        let Some(skill) = skills.iter().find(|s| s.name == skill_name) else {
763            return text.to_owned();
764        };
765
766        match std::fs::read_to_string(&skill.file_path) {
767            Ok(content) => {
768                let body = strip_frontmatter(&content)
769                    .unwrap_or_else(|_| content.clone())
770                    .trim()
771                    .to_owned();
772                let block = format!(
773                    "<skill name=\"{}\" location=\"{}\">\nReferences are relative to {}.\n\n{}\n</skill>",
774                    skill.name, skill.file_path, skill.base_dir, body
775                );
776                if args.is_empty() {
777                    block
778                } else {
779                    format!("{block}\n\n{args}")
780                }
781            }
782            Err(_) => text.to_owned(),
783        }
784    }
785
786    // -----------------------------------------------------------------
787    // Helpers
788    // -----------------------------------------------------------------
789
790    fn reserve_run_admission(self: &Arc<Self>) -> Option<RunAdmission> {
791        let mut inner = self.lock_inner();
792        if inner.is_agent_run_active {
793            return None;
794        }
795        inner.is_agent_run_active = true;
796        drop(inner);
797        Some(RunAdmission {
798            session: Arc::clone(self),
799            armed: true,
800        })
801    }
802
803    fn is_session_streaming(&self) -> bool {
804        self.lock_inner().is_agent_run_active
805    }
806
807    fn assistant_count(&self) -> usize {
808        self.agent
809            .transcript()
810            .iter()
811            .filter(|m| m.role() == "assistant")
812            .count()
813    }
814
815    fn try_model_runtime(&self) -> Option<ModelRuntime> {
816        self.model_runtime.as_deref().cloned()
817    }
818}
819
820// -----------------------------------------------------------------------
821// Free helpers
822// -----------------------------------------------------------------------
823
824fn parse_slash_command(text: &str) -> Option<(&str, &str)> {
825    if !text.starts_with('/') {
826        return None;
827    }
828    let body = &text[1..];
829    match body.find(' ') {
830        Some(idx) => Some((&body[..idx], &body[idx + 1..])),
831        None => Some((body, "")),
832    }
833}
834
835/// Whether this run outcome is followed by an `agent_end` agent event.
836///
837/// `Agent::prompt` / `Agent::continue_run` emit `agent_end` for every run that
838/// actually starts — including failed runs, whose terminal sequence is
839/// synthesized by the agent before the error returns. Only the pre-start
840/// rejections below return without emitting anything; awaiting the processed
841/// `agent_end` barrier for them would hang forever.
842fn run_emits_agent_end(run: &Result<(), pi_agent::AgentLoopError>) -> bool {
843    match run {
844        Ok(()) => true,
845        Err(pi_agent::AgentLoopError::Message(message)) => {
846            message != "agent is already running"
847                && message != "No messages to continue from"
848                && !message.starts_with("Cannot continue from message role")
849        }
850    }
851}
852
853fn is_no_model(model: &pi_ai::Model) -> bool {
854    model.provider == "unknown"
855}
856
857fn parse_images_value(value: &serde_json::Value) -> Option<Vec<ImageContent>> {
858    serde_json::from_value::<Vec<ImageContent>>(value.clone()).ok()
859}
860
861fn build_custom_agent_message(message: &CustomMessageInput) -> AgentMessage {
862    let mut payload = serde_json::Map::new();
863    payload.insert(
864        "customType".to_owned(),
865        serde_json::Value::String(message.custom_type.clone()),
866    );
867    payload.insert(
868        "content".to_owned(),
869        serde_json::to_value(&message.content).unwrap_or(serde_json::Value::Null),
870    );
871    payload.insert(
872        "display".to_owned(),
873        serde_json::Value::Bool(message.display),
874    );
875    if let Some(details) = &message.details {
876        payload.insert("details".to_owned(), details.clone());
877    }
878    payload.insert(
879        "timestamp".to_owned(),
880        serde_json::Value::Number(serde_json::Number::from(pi_agent::now_millis())),
881    );
882    AgentMessage::Custom(pi_agent::CustomAgentMessage::new("custom", payload))
883}
884
885#[cfg(test)]
886mod tests {
887    use super::*;
888    use crate::core::agent_session::{
889        AgentSessionConfig, AgentSessionEvent, ExtensionRunner, ExtensionRunnerError,
890        NullExtensionRunner,
891    };
892    use futures::future::BoxFuture;
893    use futures::stream::{self, BoxStream, StreamExt};
894    use pi_ai::{
895        AssistantContent, AssistantMessageEvent, Context, DoneReason, ErrorReason, ModelCost,
896        ModelInput, Provider, ProviderError, StopReason, StreamOptions, TextContent,
897    };
898    use std::collections::HashMap;
899    use std::fmt::Display;
900    use std::sync::Mutex as StdMutex;
901    use std::sync::atomic::{AtomicUsize, Ordering};
902    use std::sync::{MutexGuard, PoisonError};
903    use tokio::sync::{Notify, Semaphore};
904
905    type ProviderEventResult = Result<AssistantMessageEvent, ProviderError>;
906    type ProviderResponse = Vec<ProviderEventResult>;
907    type ProviderResponses = Vec<ProviderResponse>;
908    type TestResult<T = ()> = Result<T, String>;
909
910    trait TestContext<T> {
911        fn test_context(self, context: &str) -> TestResult<T>;
912    }
913
914    impl<T, E: Display> TestContext<T> for Result<T, E> {
915        fn test_context(self, context: &str) -> TestResult<T> {
916            self.map_err(|error| format!("{context}: {error}"))
917        }
918    }
919
920    fn mutex_value<T>(mutex: &StdMutex<T>) -> MutexGuard<'_, T> {
921        mutex.lock().unwrap_or_else(PoisonError::into_inner)
922    }
923
924    fn require_error<T, E>(result: Result<T, E>, context: &str) -> TestResult<E> {
925        match result {
926            Ok(_) => Err(format!("{context}: expected an error")),
927            Err(error) => Ok(error),
928        }
929    }
930
931    fn require_some<T>(value: Option<T>, context: &str) -> TestResult<T> {
932        value.ok_or_else(|| format!("{context}: expected a value"))
933    }
934
935    fn test_model() -> pi_ai::Model {
936        pi_ai::Model {
937            id: "m".to_owned(),
938            name: "m".to_owned(),
939            api: "test-api".to_owned(),
940            provider: "test-provider".to_owned(),
941            base_url: String::new(),
942            reasoning: false,
943            thinking_level_map: None,
944            input: vec![ModelInput::Text],
945            cost: ModelCost::default(),
946            context_window: 8_192,
947            max_tokens: 1_024,
948            headers: None,
949            compat: None,
950            extra: std::collections::BTreeMap::new(),
951        }
952    }
953
954    fn assistant_text(text: &str) -> AssistantMessage {
955        let mut message =
956            AssistantMessage::new("test-api", "test-provider", "m", pi_agent::now_millis());
957        message
958            .content
959            .push(AssistantContent::Text(TextContent::new(text)));
960        message.stop_reason = StopReason::Stop;
961        message
962    }
963
964    fn assistant_error(err: &str) -> AssistantMessage {
965        let mut message =
966            AssistantMessage::new("test-api", "test-provider", "m", pi_agent::now_millis());
967        message.stop_reason = StopReason::Error;
968        message.error_message = Some(err.to_owned());
969        message
970    }
971
972    fn start_event() -> AssistantMessageEvent {
973        AssistantMessageEvent::Start {
974            partial: AssistantMessage::new(
975                "test-api",
976                "test-provider",
977                "m",
978                pi_agent::now_millis(),
979            ),
980        }
981    }
982
983    fn done_ok(msg: AssistantMessage) -> AssistantMessageEvent {
984        AssistantMessageEvent::Done {
985            reason: DoneReason::Stop,
986            message: msg,
987        }
988    }
989
990    fn done_err(msg: AssistantMessage) -> AssistantMessageEvent {
991        AssistantMessageEvent::Error {
992            reason: ErrorReason::Error,
993            error: msg,
994        }
995    }
996
997    #[derive(Clone)]
998    struct SeqProvider {
999        calls: Arc<AtomicUsize>,
1000        responses: Arc<StdMutex<ProviderResponses>>,
1001    }
1002
1003    impl SeqProvider {
1004        fn new(responses: ProviderResponses) -> Self {
1005            Self {
1006                calls: Arc::new(AtomicUsize::new(0)),
1007                responses: Arc::new(StdMutex::new(responses)),
1008            }
1009        }
1010
1011        fn call_count(&self) -> usize {
1012            self.calls.load(Ordering::SeqCst)
1013        }
1014    }
1015
1016    impl Provider for SeqProvider {
1017        fn stream(
1018            &self,
1019            _model: &pi_ai::Model,
1020            _context: Context,
1021            _options: StreamOptions,
1022        ) -> BoxStream<'static, ProviderEventResult> {
1023            let idx = self.calls.fetch_add(1, Ordering::SeqCst);
1024            let events = mutex_value(&self.responses)
1025                .get(idx)
1026                .cloned()
1027                .unwrap_or_default();
1028            stream::iter(events).boxed()
1029        }
1030    }
1031
1032    fn make_session(provider: Arc<dyn Provider>) -> TestResult<Arc<AgentSession>> {
1033        let config = AgentSessionConfig::test_config(provider, test_model())
1034            .test_context("test session config")?;
1035        AgentSession::new(config).test_context("test session creation")
1036    }
1037
1038    /// Wait until the session-level run is idle (no wall-clock sleep).
1039    async fn drain(session: &Arc<AgentSession>) {
1040        session.wait_for_idle().await;
1041    }
1042
1043    /// Wrap an event into `Result<_, ProviderError>` for `Vec<Vec<_>>`.
1044    fn ok_event(e: AssistantMessageEvent) -> AssistantMessageEvent {
1045        e
1046    }
1047
1048    /// Build a response sequence where every event is wrapped in `Result::Ok`.
1049    fn sequence(events: Vec<AssistantMessageEvent>) -> ProviderResponse {
1050        events.into_iter().map(Ok).collect()
1051    }
1052
1053    /// Single-event sequence.
1054    fn one(e: AssistantMessageEvent) -> ProviderResponses {
1055        vec![sequence(vec![e])]
1056    }
1057
1058    /// Two-event sequence.
1059    fn two(a: AssistantMessageEvent, b: AssistantMessageEvent) -> ProviderResponses {
1060        vec![sequence(vec![a, b])]
1061    }
1062
1063    /// Two-call sequence (split across provider.stream calls).
1064    fn split(
1065        first: Vec<AssistantMessageEvent>,
1066        second: Vec<AssistantMessageEvent>,
1067    ) -> ProviderResponses {
1068        vec![sequence(first), sequence(second)]
1069    }
1070
1071    #[tokio::test]
1072    async fn single_prompt_records_messages() -> TestResult {
1073        let provider = Arc::new(SeqProvider::new(split(
1074            vec![
1075                ok_event(start_event()),
1076                ok_event(done_ok(assistant_text("hello"))),
1077            ],
1078            vec![],
1079        )));
1080        let session = make_session(provider)?;
1081        session
1082            .prompt("hi", PromptOptions::default())
1083            .await
1084            .test_context("single prompt")?;
1085        drain(&session).await;
1086        let messages = session.messages();
1087        let roles: Vec<&str> = messages.iter().map(AgentMessage::role).collect();
1088        assert_eq!(roles, vec!["user", "assistant"]);
1089        Ok(())
1090    }
1091
1092    #[tokio::test]
1093    async fn concurrent_prompt_without_behavior_errors() -> TestResult {
1094        let provider = Arc::new(SeqProvider::new(one(start_event())));
1095        let session = make_session(provider)?;
1096        session.mark_agent_run_active();
1097        let result = session.prompt("second", PromptOptions::default()).await;
1098        let err = require_error(result, "concurrent prompt")?;
1099        assert!(
1100            err.to_string().contains("Agent is already processing"),
1101            "{err}"
1102        );
1103        Ok(())
1104    }
1105
1106    #[tokio::test]
1107    async fn concurrent_prompts_preserve_streaming_queue_behavior() -> TestResult {
1108        let provider = Arc::new(SeqProvider::new(one(start_event())));
1109        let session = make_session(provider.clone())?;
1110        let accepted = Arc::new(AtomicUsize::new(0));
1111        session.mark_agent_run_active();
1112
1113        for behavior in [StreamingBehavior::Steer, StreamingBehavior::FollowUp] {
1114            let callback_count = Arc::clone(&accepted);
1115            session
1116                .prompt(
1117                    behavior.as_str(),
1118                    PromptOptions {
1119                        streaming_behavior: Some(behavior),
1120                        preflight_result: Some(Arc::new(move |ok| {
1121                            if ok {
1122                                callback_count.fetch_add(1, Ordering::SeqCst);
1123                            }
1124                        })),
1125                        ..PromptOptions::default()
1126                    },
1127                )
1128                .await
1129                .test_context("queued concurrent prompt")?;
1130        }
1131
1132        assert_eq!(accepted.load(Ordering::SeqCst), 2);
1133        assert_eq!(session.pending_message_count(), 2);
1134        assert_eq!(provider.call_count(), 0);
1135        Ok(())
1136    }
1137
1138    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1139    async fn concurrent_prompts_admit_exactly_one_run() -> TestResult {
1140        let provider = Arc::new(SeqProvider::new(vec![
1141            sequence(vec![start_event(), done_ok(assistant_text("first"))]),
1142            sequence(vec![start_event(), done_ok(assistant_text("second"))]),
1143        ]));
1144        let session = make_session(provider.clone())?;
1145        let accepted = Arc::new(AtomicUsize::new(0));
1146        let first_entered = Arc::new(Semaphore::new(0));
1147        let first_gate = Arc::new(std::sync::Barrier::new(2));
1148
1149        let first_session = Arc::clone(&session);
1150        let first_accepted = Arc::clone(&accepted);
1151        let first_entered_callback = Arc::clone(&first_entered);
1152        let first_gate_callback = Arc::clone(&first_gate);
1153        let first = tokio::spawn(async move {
1154            first_session
1155                .prompt(
1156                    "first",
1157                    PromptOptions {
1158                        preflight_result: Some(Arc::new(move |ok| {
1159                            if ok {
1160                                first_accepted.fetch_add(1, Ordering::SeqCst);
1161                                first_entered_callback.add_permits(1);
1162                                first_gate_callback.wait();
1163                            }
1164                        })),
1165                        ..PromptOptions::default()
1166                    },
1167                )
1168                .await
1169        });
1170
1171        first_entered
1172            .acquire()
1173            .await
1174            .test_context("first preflight signal")?
1175            .forget();
1176
1177        let second_accepted = Arc::clone(&accepted);
1178        let second = session
1179            .prompt(
1180                "second",
1181                PromptOptions {
1182                    preflight_result: Some(Arc::new(move |ok| {
1183                        if ok {
1184                            second_accepted.fetch_add(1, Ordering::SeqCst);
1185                        }
1186                    })),
1187                    ..PromptOptions::default()
1188                },
1189            )
1190            .await;
1191
1192        first_gate.wait();
1193        first
1194            .await
1195            .test_context("first prompt task")?
1196            .test_context("first prompt")?;
1197        let err = require_error(second, "second prompt admission")?;
1198        assert!(
1199            err.to_string().contains("Agent is already processing"),
1200            "{err}"
1201        );
1202        assert_eq!(accepted.load(Ordering::SeqCst), 1);
1203        assert_eq!(provider.call_count(), 1);
1204        Ok(())
1205    }
1206
1207    #[tokio::test]
1208    async fn panicking_preflight_callback_releases_run_admission() -> TestResult {
1209        let provider = Arc::new(SeqProvider::new(two(
1210            start_event(),
1211            done_ok(assistant_text("after panic")),
1212        )));
1213        let session = make_session(provider.clone())?;
1214        let panicking_session = Arc::clone(&session);
1215        let panicking = tokio::spawn(async move {
1216            panicking_session
1217                .prompt(
1218                    "panic",
1219                    PromptOptions {
1220                        preflight_result: Some(Arc::new(|ok| {
1221                            assert!(!ok, "intentional accepted-callback panic");
1222                        })),
1223                        ..PromptOptions::default()
1224                    },
1225                )
1226                .await
1227        });
1228
1229        let join_error = require_error(panicking.await, "panicking prompt task")?;
1230        assert!(join_error.is_panic());
1231        session
1232            .prompt("after panic", PromptOptions::default())
1233            .await
1234            .test_context("prompt after callback panic")?;
1235        assert_eq!(provider.call_count(), 1);
1236        Ok(())
1237    }
1238
1239    #[tokio::test]
1240    async fn no_model_error() -> TestResult {
1241        let provider = Arc::new(SeqProvider::new(one(start_event())));
1242        let mut config =
1243            AgentSessionConfig::test_config(provider, test_model()).test_context("test config")?;
1244        config.model = None;
1245        let session = AgentSession::new(config).test_context("session")?;
1246        let result = session.prompt("hi", PromptOptions::default()).await;
1247        let err = require_error(result, "no-model prompt")?;
1248        assert!(err.to_string().contains("No model selected"), "{err}");
1249        assert!(!session.is_session_streaming());
1250        Ok(())
1251    }
1252
1253    #[tokio::test]
1254    async fn retry_transient_then_success() -> TestResult {
1255        let provider = Arc::new(SeqProvider::new(split(
1256            vec![
1257                ok_event(start_event()),
1258                ok_event(done_err(assistant_error("overloaded_error"))),
1259            ],
1260            vec![
1261                ok_event(start_event()),
1262                ok_event(done_ok(assistant_text("recovered"))),
1263            ],
1264        )));
1265        let session = make_session(provider.clone())?;
1266        let events = Arc::new(StdMutex::new(Vec::<String>::new()));
1267        let settled = Arc::new(AtomicUsize::new(0));
1268        let ev = events.clone();
1269        let s = settled.clone();
1270        let _u = session.subscribe(move |event| match event {
1271            AgentSessionEvent::AutoRetryStart { attempt, .. } => {
1272                mutex_value(&ev).push(format!("start:{attempt}"));
1273            }
1274            AgentSessionEvent::AutoRetryEnd { success, .. } => {
1275                mutex_value(&ev).push(format!("end:{success}"));
1276            }
1277            AgentSessionEvent::AgentSettled => {
1278                s.fetch_add(1, Ordering::SeqCst);
1279            }
1280            _ => {}
1281        });
1282        session
1283            .prompt("test", PromptOptions::default())
1284            .await
1285            .test_context("prompt")?;
1286        drain(&session).await;
1287        let ev = mutex_value(&events).clone();
1288        assert_eq!(
1289            provider.call_count(),
1290            2,
1291            "one failure + one recovery stream"
1292        );
1293        assert_eq!(
1294            ev,
1295            vec!["start:1".to_owned(), "end:true".to_owned()],
1296            "retry lifecycle order: {ev:?}"
1297        );
1298        assert_eq!(settled.load(Ordering::SeqCst), 1);
1299        assert_eq!(session.retry_attempt(), 0);
1300        assert!(!session.agent.state().is_streaming);
1301        Ok(())
1302    }
1303
1304    #[tokio::test]
1305    async fn retry_disabled_no_retry() -> TestResult {
1306        let provider = Arc::new(SeqProvider::new(two(
1307            start_event(),
1308            done_err(assistant_error("overloaded_error")),
1309        )));
1310        let session = make_session(provider.clone())?;
1311        session.set_auto_retry_enabled(false);
1312        let events = Arc::new(StdMutex::new(Vec::<String>::new()));
1313        let settled = Arc::new(AtomicUsize::new(0));
1314        let ev = events.clone();
1315        let s = settled.clone();
1316        let _u = session.subscribe(move |e| match e {
1317            AgentSessionEvent::AutoRetryStart { .. } => {
1318                mutex_value(&ev).push("start".to_owned());
1319            }
1320            AgentSessionEvent::AutoRetryEnd { .. } => {
1321                mutex_value(&ev).push("end".to_owned());
1322            }
1323            AgentSessionEvent::AgentSettled => {
1324                s.fetch_add(1, Ordering::SeqCst);
1325            }
1326            _ => {}
1327        });
1328        session
1329            .prompt("test", PromptOptions::default())
1330            .await
1331            .test_context("prompt")?;
1332        drain(&session).await;
1333        let ev = mutex_value(&events).clone();
1334        assert!(
1335            ev.is_empty(),
1336            "disabled retry must not emit auto_retry events: {ev:?}"
1337        );
1338        assert_eq!(
1339            provider.call_count(),
1340            1,
1341            "disabled retry must not re-invoke the provider"
1342        );
1343        assert_eq!(settled.load(Ordering::SeqCst), 1);
1344        assert_eq!(session.retry_attempt(), 0);
1345        assert!(!session.auto_retry_enabled());
1346        Ok(())
1347    }
1348
1349    #[tokio::test]
1350    async fn non_retryable_error_no_retry() -> TestResult {
1351        let provider = Arc::new(SeqProvider::new(two(
1352            start_event(),
1353            done_err(assistant_error("invalid_api_key")),
1354        )));
1355        let session = make_session(provider.clone())?;
1356        let events = Arc::new(StdMutex::new(Vec::<String>::new()));
1357        let settled = Arc::new(AtomicUsize::new(0));
1358        let ev = events.clone();
1359        let s = settled.clone();
1360        let _u = session.subscribe(move |e| match e {
1361            AgentSessionEvent::AutoRetryStart { .. } => {
1362                mutex_value(&ev).push("start".to_owned());
1363            }
1364            AgentSessionEvent::AutoRetryEnd { .. } => {
1365                mutex_value(&ev).push("end".to_owned());
1366            }
1367            AgentSessionEvent::AgentSettled => {
1368                s.fetch_add(1, Ordering::SeqCst);
1369            }
1370            _ => {}
1371        });
1372        session
1373            .prompt("test", PromptOptions::default())
1374            .await
1375            .test_context("prompt")?;
1376        drain(&session).await;
1377        let ev = mutex_value(&events).clone();
1378        assert!(
1379            ev.is_empty(),
1380            "auth error must not emit auto_retry events: {ev:?}"
1381        );
1382        assert_eq!(provider.call_count(), 1);
1383        assert_eq!(settled.load(Ordering::SeqCst), 1);
1384        assert_eq!(session.retry_attempt(), 0);
1385        Ok(())
1386    }
1387
1388    #[tokio::test]
1389    async fn single_settled_after_prompt() -> TestResult {
1390        let provider = Arc::new(SeqProvider::new(two(
1391            start_event(),
1392            done_ok(assistant_text("ok")),
1393        )));
1394        let session = make_session(provider)?;
1395        let count = Arc::new(AtomicUsize::new(0));
1396        let c = count.clone();
1397        let _u = session.subscribe(move |e| {
1398            if matches!(e, AgentSessionEvent::AgentSettled) {
1399                c.fetch_add(1, Ordering::SeqCst);
1400            }
1401        });
1402        session
1403            .prompt("hi", PromptOptions::default())
1404            .await
1405            .test_context("prompt")?;
1406        drain(&session).await;
1407        assert_eq!(count.load(Ordering::SeqCst), 1);
1408        Ok(())
1409    }
1410
1411    #[tokio::test]
1412    async fn settled_waits_for_agent_end_extension_processing() -> TestResult {
1413        let gate = Arc::new(Semaphore::new(0));
1414        let entered = Arc::new(Notify::new());
1415        let runner = Arc::new(TestRunner {
1416            agent_end_gate: Some(Arc::clone(&gate)),
1417            agent_end_entered: Some(Arc::clone(&entered)),
1418            ..TestRunner::default()
1419        });
1420        let provider = Arc::new(SeqProvider::new(two(
1421            start_event(),
1422            done_ok(assistant_text("ok")),
1423        )));
1424        let mut config =
1425            AgentSessionConfig::test_config(provider, test_model()).test_context("test config")?;
1426        config.extension_runner = Some(runner as Arc<dyn ExtensionRunner>);
1427        let session = AgentSession::new(config).test_context("session")?;
1428        let settled = Arc::new(AtomicUsize::new(0));
1429        let settled_for_listener = Arc::clone(&settled);
1430        let _unsubscribe = session.subscribe(move |event| {
1431            if matches!(event, AgentSessionEvent::AgentSettled) {
1432                settled_for_listener.fetch_add(1, Ordering::SeqCst);
1433            }
1434        });
1435
1436        let entered_wait = entered.notified();
1437        let session_for_prompt = Arc::clone(&session);
1438        let prompt = tokio::spawn(async move {
1439            session_for_prompt
1440                .prompt("hi", PromptOptions::default())
1441                .await
1442        });
1443        entered_wait.await;
1444        assert_eq!(settled.load(Ordering::SeqCst), 0);
1445
1446        gate.add_permits(1);
1447        prompt
1448            .await
1449            .test_context("joining gated prompt")?
1450            .test_context("gated prompt")?;
1451        assert_eq!(settled.load(Ordering::SeqCst), 1);
1452        Ok(())
1453    }
1454
1455    #[tokio::test]
1456    async fn disconnect_cancels_agent_end_barrier() -> TestResult {
1457        let provider = Arc::new(SeqProvider::new(one(start_event())));
1458        let session = make_session(provider)?;
1459        let before = session.processed_agent_end_count();
1460        let session_for_wait = Arc::clone(&session);
1461        let waiter =
1462            tokio::spawn(
1463                async move { session_for_wait.wait_for_processed_agent_end(before).await },
1464            );
1465        tokio::task::yield_now().await;
1466
1467        session.disconnect_from_agent();
1468        let wait_completed = waiter.await.test_context("joining agent-end waiter")?;
1469        assert!(!wait_completed);
1470        session.reconnect_to_agent();
1471        session.dispose().await;
1472        Ok(())
1473    }
1474
1475    #[tokio::test]
1476    async fn queue_steer_and_follow_up() -> TestResult {
1477        let provider = Arc::new(SeqProvider::new(one(start_event())));
1478        let session = make_session(provider)?;
1479        session.queue_steer("a", Vec::new());
1480        assert_eq!(session.pending_message_count(), 1);
1481        session.queue_follow_up("b", Vec::new());
1482        assert_eq!(session.pending_message_count(), 2);
1483        session.clear_queue();
1484        assert_eq!(session.pending_message_count(), 0);
1485        Ok(())
1486    }
1487
1488    #[tokio::test]
1489    async fn prompt_preflight_flushes_bash_before_validation() -> TestResult {
1490        let provider = Arc::new(SeqProvider::new(two(
1491            start_event(),
1492            done_ok(assistant_text("unused")),
1493        )));
1494        let mut config =
1495            AgentSessionConfig::test_config(provider, test_model()).test_context("test config")?;
1496        config.model = None;
1497        let session = AgentSession::new(config).test_context("session")?;
1498        session.lock_inner().pending_bash_messages.push(
1499            crate::core::messages::BashExecutionMessage::from_fields(
1500                crate::core::messages::BashExecutionFields {
1501                    command: "printf pending".to_owned(),
1502                    output: "pending".to_owned(),
1503                    exit_code: Some(0),
1504                    cancelled: false,
1505                    truncated: false,
1506                    full_output_path: None,
1507                    timestamp: 1,
1508                    exclude_from_context: None,
1509                },
1510            ),
1511        );
1512
1513        assert!(session.has_pending_bash_messages());
1514        let result = session.prompt("x", PromptOptions::default()).await;
1515
1516        assert!(matches!(result, Err(PromptError::Message(_))));
1517        assert!(!session.has_pending_bash_messages());
1518        Ok(())
1519    }
1520
1521    #[tokio::test]
1522    async fn prompt_flush_failure_retains_bash_message_for_retry() -> TestResult {
1523        let dir = tempfile::tempdir().test_context("tempdir")?;
1524        let mut manager = crate::core::sessions::SessionManager::create(
1525            dir.path().to_string_lossy().as_ref(),
1526            Some(dir.path().to_string_lossy().as_ref()),
1527            None,
1528        )
1529        .test_context("session manager")?;
1530        manager
1531            .append_message(&pi_agent::user_text("hi", std::iter::empty()))
1532            .test_context("append user")?;
1533        let assistant = AgentMessage::Llm(Box::new(pi_ai::Message::Assistant(assistant_text(
1534            "answer",
1535        ))));
1536        manager
1537            .append_message(&assistant)
1538            .test_context("append assistant")?;
1539        let session_file = std::path::PathBuf::from(
1540            manager
1541                .get_session_file()
1542                .ok_or_else(|| "missing session file".to_owned())?,
1543        );
1544        let backup = dir.path().join("session-backup.jsonl");
1545        std::fs::rename(&session_file, &backup).test_context("move session aside")?;
1546        std::fs::create_dir(&session_file).test_context("block append path")?;
1547
1548        let provider = Arc::new(SeqProvider::new(one(start_event())));
1549        let mut config =
1550            AgentSessionConfig::test_config(provider, test_model()).test_context("test config")?;
1551        config.model = None;
1552        config.session_manager = manager;
1553        let session = AgentSession::new(config).test_context("session")?;
1554        let bash_message = |command: &str, timestamp: i64| {
1555            crate::core::messages::BashExecutionMessage::from_fields(
1556                crate::core::messages::BashExecutionFields {
1557                    command: command.to_owned(),
1558                    output: command.to_owned(),
1559                    exit_code: Some(0),
1560                    cancelled: false,
1561                    truncated: false,
1562                    full_output_path: None,
1563                    timestamp,
1564                    exclude_from_context: None,
1565                },
1566            )
1567        };
1568        session.lock_inner().pending_bash_messages.extend([
1569            bash_message("printf first", 1),
1570            bash_message("printf second", 2),
1571        ]);
1572
1573        let err = require_error(
1574            session.prompt("x", PromptOptions::default()).await,
1575            "persistence failure",
1576        )?;
1577        assert!(matches!(err, PromptError::Session(_)));
1578        assert_eq!(session.lock_inner().pending_bash_messages.len(), 2);
1579
1580        std::fs::remove_dir(&session_file).test_context("remove append blocker")?;
1581        std::fs::rename(&backup, &session_file).test_context("restore session")?;
1582        session
1583            .flush_pending_bash_messages()
1584            .await
1585            .test_context("retry flush")?;
1586        assert!(!session.has_pending_bash_messages());
1587        let persisted =
1588            std::fs::read_to_string(&session_file).test_context("read persisted session")?;
1589        let first = persisted
1590            .find("printf first")
1591            .ok_or_else(|| "missing first bash".to_owned())?;
1592        let second = persisted
1593            .find("printf second")
1594            .ok_or_else(|| "missing second bash".to_owned())?;
1595        assert!(
1596            first < second,
1597            "retried bash messages must preserve queue order"
1598        );
1599        assert_eq!(persisted.matches("\"role\":\"bashExecution\"").count(), 2);
1600        Ok(())
1601    }
1602
1603    #[tokio::test]
1604    async fn message_end_disk_failure_returns_typed_prompt_error() -> TestResult {
1605        let dir = tempfile::tempdir().test_context("tempdir")?;
1606        let mut manager = crate::core::sessions::SessionManager::create(
1607            dir.path().to_string_lossy().as_ref(),
1608            Some(dir.path().to_string_lossy().as_ref()),
1609            None,
1610        )
1611        .test_context("session manager")?;
1612        manager
1613            .append_message(&pi_agent::user_text("existing", std::iter::empty()))
1614            .test_context("append existing user")?;
1615        manager
1616            .append_message(&AgentMessage::Llm(Box::new(pi_ai::Message::Assistant(
1617                assistant_text("existing answer"),
1618            ))))
1619            .test_context("append existing assistant")?;
1620        let before_count = manager.get_entries().len();
1621        let session_file = std::path::PathBuf::from(
1622            manager
1623                .get_session_file()
1624                .ok_or_else(|| "missing session file".to_owned())?,
1625        );
1626        let backup = dir.path().join("message-end-backup.jsonl");
1627        std::fs::rename(&session_file, &backup).test_context("move session aside")?;
1628        std::fs::create_dir(&session_file).test_context("block append path")?;
1629
1630        let provider = Arc::new(SeqProvider::new(two(
1631            start_event(),
1632            done_ok(assistant_text("new answer")),
1633        )));
1634        let mut config =
1635            AgentSessionConfig::test_config(provider, test_model()).test_context("test config")?;
1636        config.session_manager = manager;
1637        let session = AgentSession::new(config).test_context("session")?;
1638
1639        let err = require_error(
1640            session
1641                .prompt("new question", PromptOptions::default())
1642                .await,
1643            "message-end persistence failure",
1644        )?;
1645        assert!(matches!(err, PromptError::Session(_)));
1646        assert_eq!(
1647            session.session_manager.lock().await.get_entries().len(),
1648            before_count,
1649            "failed append must not advance the in-memory tree"
1650        );
1651
1652        std::fs::remove_dir(&session_file).test_context("remove append blocker")?;
1653        std::fs::rename(&backup, &session_file).test_context("restore session")?;
1654        Ok(())
1655    }
1656
1657    #[tokio::test]
1658    async fn handle_post_settles_after_retry_path_without_queue() -> TestResult {
1659        // Observable post-run contract: retry → (no queue) → exactly one settle.
1660        // A successful recovery leaves retry_attempt at 0 and no pending queue.
1661        let provider = Arc::new(SeqProvider::new(split(
1662            vec![
1663                ok_event(start_event()),
1664                ok_event(done_err(assistant_error("overloaded_error"))),
1665            ],
1666            vec![
1667                ok_event(start_event()),
1668                ok_event(done_ok(assistant_text("ok"))),
1669            ],
1670        )));
1671        let session = make_session(provider.clone())?;
1672        let order = Arc::new(StdMutex::new(Vec::<String>::new()));
1673        let o = order.clone();
1674        let _u = session.subscribe(move |e| match e {
1675            AgentSessionEvent::AutoRetryStart { .. } => {
1676                mutex_value(&o).push("retry".to_owned());
1677            }
1678            AgentSessionEvent::AutoRetryEnd { success, .. } => {
1679                mutex_value(&o).push(format!("retry_end:{success}"));
1680            }
1681            AgentSessionEvent::AgentSettled => {
1682                mutex_value(&o).push("settled".to_owned());
1683            }
1684            _ => {}
1685        });
1686        session
1687            .prompt("test", PromptOptions::default())
1688            .await
1689            .test_context("prompt")?;
1690        drain(&session).await;
1691        let order = mutex_value(&order).clone();
1692        assert_eq!(
1693            order,
1694            vec![
1695                "retry".to_owned(),
1696                "retry_end:true".to_owned(),
1697                "settled".to_owned()
1698            ],
1699            "post-run order: {order:?}"
1700        );
1701        assert_eq!(provider.call_count(), 2);
1702        assert_eq!(session.pending_message_count(), 0);
1703        assert_eq!(session.retry_attempt(), 0);
1704        let last = require_some(session.agent.last_assistant(), "last assistant after retry")?;
1705        assert_eq!(last.stop_reason, StopReason::Stop);
1706        Ok(())
1707    }
1708
1709    #[tokio::test]
1710    async fn run_agent_prompt_flushes_bash_before_settled() -> TestResult {
1711        // Smoke: completion settles exactly once after the prompt loop.
1712        let provider = Arc::new(SeqProvider::new(two(
1713            start_event(),
1714            done_ok(assistant_text("ok")),
1715        )));
1716        let session = make_session(provider)?;
1717        let count = Arc::new(AtomicUsize::new(0));
1718        let c = count.clone();
1719        let _u = session.subscribe(move |e| {
1720            if matches!(e, AgentSessionEvent::AgentSettled) {
1721                c.fetch_add(1, Ordering::SeqCst);
1722            }
1723        });
1724        session
1725            .prompt("hi", PromptOptions::default())
1726            .await
1727            .test_context("prompt")?;
1728        drain(&session).await;
1729        assert_eq!(count.load(Ordering::SeqCst), 1);
1730        Ok(())
1731    }
1732
1733    #[tokio::test]
1734    async fn null_runner_command_passes_through() -> TestResult {
1735        let provider = Arc::new(SeqProvider::new(two(
1736            start_event(),
1737            done_ok(assistant_text("ok")),
1738        )));
1739        let session = make_session(provider)?;
1740        session
1741            .prompt("/nonexistent hello", PromptOptions::default())
1742            .await
1743            .test_context("prompt")?;
1744        drain(&session).await;
1745        let messages = session.messages();
1746        let roles: Vec<&str> = messages.iter().map(AgentMessage::role).collect();
1747        assert_eq!(roles, vec!["user", "assistant"]);
1748        Ok(())
1749    }
1750
1751    #[derive(Default)]
1752    struct TestRunner {
1753        commands: Vec<String>,
1754        runs: Arc<StdMutex<Vec<String>>>,
1755        agent_end_gate: Option<Arc<Semaphore>>,
1756        agent_end_entered: Option<Arc<Notify>>,
1757        before_start_gate: Option<Arc<Semaphore>>,
1758        before_start_entered: Option<Arc<Notify>>,
1759    }
1760
1761    impl ExtensionRunner for TestRunner {
1762        fn has_handlers(&self, event: &str) -> bool {
1763            event == "agent_end"
1764                && (self.agent_end_gate.is_some() || self.agent_end_entered.is_some())
1765        }
1766        fn emit(
1767            &self,
1768            event: AgentSessionEvent,
1769        ) -> BoxFuture<
1770            '_,
1771            Result<Option<crate::core::agent_session::CancelResult>, ExtensionRunnerError>,
1772        > {
1773            let gate = self.agent_end_gate.clone();
1774            let entered = self.agent_end_entered.clone();
1775            Box::pin(async move {
1776                if matches!(event, AgentSessionEvent::AgentEnd { .. }) {
1777                    if let Some(entered) = entered {
1778                        entered.notify_one();
1779                    }
1780                    if let Some(gate) = gate
1781                        && let Ok(permit) = gate.acquire_owned().await
1782                    {
1783                        permit.forget();
1784                    }
1785                }
1786                Ok(None)
1787            })
1788        }
1789        fn emit_message_end(
1790            &self,
1791            _m: AgentMessage,
1792        ) -> BoxFuture<'_, Result<Option<AgentMessage>, ExtensionRunnerError>> {
1793            Box::pin(async { Ok(None) })
1794        }
1795        fn emit_tool_call(
1796            &self,
1797            _: &str,
1798            _: &str,
1799            _: serde_json::Map<String, serde_json::Value>,
1800        ) -> BoxFuture<'_, Result<Option<pi_agent::BeforeToolCallResult>, ExtensionRunnerError>>
1801        {
1802            Box::pin(async { Ok(None) })
1803        }
1804        fn emit_tool_result(
1805            &self,
1806            _: &str,
1807            _: &str,
1808            _: serde_json::Map<String, serde_json::Value>,
1809            _: Vec<pi_ai::ToolResultContent>,
1810            _: serde_json::Value,
1811            _: bool,
1812        ) -> BoxFuture<'_, Result<Option<pi_agent::AfterToolCallResult>, ExtensionRunnerError>>
1813        {
1814            Box::pin(async { Ok(None) })
1815        }
1816        fn emit_input(
1817            &self,
1818            _: &str,
1819            _: Option<serde_json::Value>,
1820            _: &str,
1821            _: Option<&str>,
1822        ) -> BoxFuture<
1823            '_,
1824            Result<crate::core::agent_session::InputTransformResult, ExtensionRunnerError>,
1825        > {
1826            Box::pin(async { Ok(crate::core::agent_session::InputTransformResult::default()) })
1827        }
1828        fn emit_before_agent_start(
1829            &self,
1830            _: &str,
1831            _: Option<serde_json::Value>,
1832        ) -> BoxFuture<
1833            '_,
1834            Result<
1835                Option<crate::core::agent_session::BeforeAgentStartResult>,
1836                ExtensionRunnerError,
1837            >,
1838        > {
1839            let gate = self.before_start_gate.clone();
1840            let entered = self.before_start_entered.clone();
1841            Box::pin(async move {
1842                if let Some(entered) = entered {
1843                    entered.notify_one();
1844                }
1845                if let Some(gate) = gate
1846                    && let Ok(permit) = gate.acquire_owned().await
1847                {
1848                    permit.forget();
1849                }
1850                Ok(None)
1851            })
1852        }
1853        fn emit_resources_discover(
1854            &self,
1855            _: &str,
1856            _: &str,
1857        ) -> BoxFuture<
1858            '_,
1859            Result<crate::core::resources::ResourceExtensionPaths, ExtensionRunnerError>,
1860        > {
1861            Box::pin(async { Ok(crate::core::resources::ResourceExtensionPaths::default()) })
1862        }
1863        fn get_registered_commands(&self) -> Vec<String> {
1864            self.commands.clone()
1865        }
1866        fn execute_command<'a>(
1867            &'a self,
1868            name: &'a str,
1869            args: &'a str,
1870        ) -> BoxFuture<'a, Result<bool, ExtensionRunnerError>> {
1871            mutex_value(&self.runs).push(format!("{name}:{args}"));
1872            Box::pin(async { Ok(true) })
1873        }
1874        fn get_all_registered_tools(&self) -> HashMap<String, Arc<dyn pi_agent::AgentTool>> {
1875            HashMap::new()
1876        }
1877        fn get_flag_values(&self) -> HashMap<String, serde_json::Value> {
1878            HashMap::new()
1879        }
1880        fn invalidate(&self) {}
1881        fn emit_error(&self, _: String) {}
1882    }
1883
1884    #[tokio::test]
1885    async fn extension_command_dispatched_idle() -> TestResult {
1886        let runner = Arc::new(TestRunner {
1887            commands: vec!["testcmd".to_owned()],
1888            runs: Arc::new(StdMutex::new(Vec::new())),
1889            ..TestRunner::default()
1890        });
1891        let provider = Arc::new(SeqProvider::new(two(
1892            start_event(),
1893            done_ok(assistant_text("queued")),
1894        )));
1895        let mut config =
1896            AgentSessionConfig::test_config(provider, test_model()).test_context("test config")?;
1897        config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
1898        let session = AgentSession::new(config).test_context("session")?;
1899
1900        session
1901            .prompt("/testcmd hello world", PromptOptions::default())
1902            .await
1903            .test_context("prompt")?;
1904
1905        let runs = mutex_value(&runner.runs).clone();
1906        assert_eq!(runs, vec!["testcmd:hello world"]);
1907        assert!(session.messages().is_empty());
1908        Ok(())
1909    }
1910
1911    #[tokio::test]
1912    async fn steer_extension_command_rejected() -> TestResult {
1913        let runner = Arc::new(TestRunner {
1914            commands: vec!["testcmd".to_owned()],
1915            runs: Arc::new(StdMutex::new(Vec::new())),
1916            ..TestRunner::default()
1917        });
1918        let provider = Arc::new(SeqProvider::new(one(start_event())));
1919        let mut config = AgentSessionConfig::test_config(provider, test_model())
1920            .map_err(|error| format!("test config failed: {error}"))?;
1921        config.extension_runner = Some(runner as Arc<dyn ExtensionRunner>);
1922        let session = AgentSession::new(config)
1923            .map_err(|error| format!("session creation failed: {error}"))?;
1924
1925        let err = match session.steer("/testcmd x", Vec::new()) {
1926            Ok(()) => return Err("extension command unexpectedly queued".to_owned()),
1927            Err(error) => error,
1928        };
1929        assert!(
1930            err.to_string()
1931                .contains("Extension command \"/testcmd\" cannot be queued"),
1932            "{err}"
1933        );
1934        Ok(())
1935    }
1936
1937    #[tokio::test]
1938    async fn retry_exhaust_emits_failure() -> TestResult {
1939        // max_retries default 3 → 4 provider streams (initial + 3 retries), then end:false.
1940        let provider = Arc::new(SeqProvider::new(vec![
1941            sequence(vec![
1942                ok_event(start_event()),
1943                ok_event(done_err(assistant_error("overloaded_error"))),
1944            ]),
1945            sequence(vec![
1946                ok_event(start_event()),
1947                ok_event(done_err(assistant_error("overloaded_error"))),
1948            ]),
1949            sequence(vec![
1950                ok_event(start_event()),
1951                ok_event(done_err(assistant_error("overloaded_error"))),
1952            ]),
1953            sequence(vec![
1954                ok_event(start_event()),
1955                ok_event(done_err(assistant_error("overloaded_error"))),
1956            ]),
1957        ]));
1958        let session = make_session(provider.clone())?;
1959        let events = Arc::new(StdMutex::new(Vec::<String>::new()));
1960        let settled = Arc::new(AtomicUsize::new(0));
1961        let ev = events.clone();
1962        let s = settled.clone();
1963        let _u = session.subscribe(move |event| match event {
1964            AgentSessionEvent::AutoRetryStart { attempt, .. } => {
1965                mutex_value(&ev).push(format!("start:{attempt}"));
1966            }
1967            AgentSessionEvent::AutoRetryEnd {
1968                success, attempt, ..
1969            } => {
1970                mutex_value(&ev).push(format!("end:{success}:{attempt}"));
1971            }
1972            AgentSessionEvent::AgentSettled => {
1973                s.fetch_add(1, Ordering::SeqCst);
1974            }
1975            _ => {}
1976        });
1977        session
1978            .prompt("test", PromptOptions::default())
1979            .await
1980            .test_context("prompt")?;
1981        drain(&session).await;
1982        let ev = mutex_value(&events).clone();
1983        assert_eq!(
1984            provider.call_count(),
1985            4,
1986            "initial + max_retries=3 must exhaust without phantom calls"
1987        );
1988        assert_eq!(
1989            ev,
1990            vec![
1991                "start:1".to_owned(),
1992                "start:2".to_owned(),
1993                "start:3".to_owned(),
1994                "end:false:3".to_owned(),
1995            ],
1996            "exhaust lifecycle: {ev:?}"
1997        );
1998        assert_eq!(settled.load(Ordering::SeqCst), 1);
1999        assert_eq!(session.retry_attempt(), 0);
2000        assert!(!session.agent.state().is_streaming);
2001        Ok(())
2002    }
2003
2004    #[tokio::test]
2005    async fn abort_retry_during_sleep() -> TestResult {
2006        let provider = Arc::new(SeqProvider::new(two(
2007            start_event(),
2008            done_err(assistant_error("overloaded_error")),
2009        )));
2010        let session = make_session(provider.clone())?;
2011        {
2012            let mut inner = session.lock_inner();
2013            inner.max_retries = 3;
2014        }
2015        let events = Arc::new(StdMutex::new(Vec::<String>::new()));
2016        let ev = events.clone();
2017        let session_for_abort = Arc::clone(&session);
2018        let _u = session.subscribe(move |event| match event {
2019            AgentSessionEvent::AutoRetryStart { attempt, .. } => {
2020                mutex_value(&ev).push(format!("start:{attempt}"));
2021                let session = Arc::clone(&session_for_abort);
2022                tokio::spawn(async move {
2023                    tokio::task::yield_now().await;
2024                    session.abort_retry();
2025                });
2026            }
2027            AgentSessionEvent::AutoRetryEnd {
2028                success,
2029                final_error,
2030                ..
2031            } => {
2032                mutex_value(&ev).push(format!(
2033                    "end:{success}:{}",
2034                    final_error.as_deref().unwrap_or("")
2035                ));
2036            }
2037            AgentSessionEvent::AgentSettled => {
2038                mutex_value(&ev).push("settled".to_owned());
2039            }
2040            _ => {}
2041        });
2042        session
2043            .prompt("test", PromptOptions::default())
2044            .await
2045            .test_context("prompt")?;
2046        drain(&session).await;
2047        let ev = mutex_value(&events).clone();
2048        assert_eq!(
2049            ev,
2050            vec![
2051                "start:1".to_owned(),
2052                "end:false:Retry cancelled".to_owned(),
2053                "settled".to_owned(),
2054            ],
2055            "abort lifecycle: {ev:?}"
2056        );
2057        assert_eq!(
2058            provider.call_count(),
2059            1,
2060            "abort during sleep must not start another provider stream"
2061        );
2062        assert_eq!(session.retry_attempt(), 0);
2063        Ok(())
2064    }
2065
2066    #[tokio::test]
2067    async fn retry_then_tool_loop_keeps_prompt_open() -> TestResult {
2068        let provider = Arc::new(SeqProvider::new(split(
2069            vec![
2070                ok_event(start_event()),
2071                ok_event(done_err(assistant_error("overloaded_error"))),
2072            ],
2073            vec![
2074                ok_event(start_event()),
2075                ok_event(done_ok(assistant_text("recovered"))),
2076            ],
2077        )));
2078        let session = make_session(provider)?;
2079        session
2080            .prompt("test", PromptOptions::default())
2081            .await
2082            .test_context("prompt")?;
2083        drain(&session).await;
2084        assert!(!session.agent.state().is_streaming);
2085        Ok(())
2086    }
2087
2088    /// Session whose next append fails (session file path blocked by a dir).
2089    fn blocked_append_session(
2090        provider: Arc<dyn Provider>,
2091        dir: &tempfile::TempDir,
2092    ) -> TestResult<Arc<AgentSession>> {
2093        let mut manager = crate::core::sessions::SessionManager::create(
2094            dir.path().to_string_lossy().as_ref(),
2095            Some(dir.path().to_string_lossy().as_ref()),
2096            None,
2097        )
2098        .test_context("session manager")?;
2099        manager
2100            .append_message(&pi_agent::user_text("seed", std::iter::empty()))
2101            .test_context("seed user append")?;
2102        // Persistence is lazy until an assistant entry exists; materialize the
2103        // file so the directory blocker makes every later append fail.
2104        manager
2105            .append_message(&AgentMessage::Llm(Box::new(pi_ai::Message::Assistant(
2106                assistant_text("seed answer"),
2107            ))))
2108            .test_context("seed assistant append")?;
2109        let session_file = std::path::PathBuf::from(
2110            manager
2111                .get_session_file()
2112                .ok_or_else(|| "missing session file".to_owned())?,
2113        );
2114        std::fs::remove_file(&session_file).test_context("remove session file")?;
2115        std::fs::create_dir(&session_file).test_context("block append path")?;
2116        let mut config =
2117            AgentSessionConfig::test_config(provider, test_model()).test_context("test config")?;
2118        config.session_manager = manager;
2119        AgentSession::new(config).test_context("session")
2120    }
2121
2122    #[tokio::test]
2123    async fn provider_error_run_observes_agent_end_before_settled() -> TestResult {
2124        let provider = Arc::new(SeqProvider::new(vec![vec![Err(
2125            pi_ai::ProviderError::new("stream exploded"),
2126        )]]));
2127        let session = make_session(provider)?;
2128        let order = Arc::new(StdMutex::new(Vec::new()));
2129        let order_clone = Arc::clone(&order);
2130        let _unsub = session.subscribe(move |event| {
2131            mutex_value(&order_clone).push(event.type_name().to_owned());
2132        });
2133
2134        // The provider failure surfaces as an error assistant terminal; the
2135        // regression under test is event ordering, not the prompt result.
2136        let _ = session.prompt("hi", PromptOptions::default()).await;
2137
2138        let observed = mutex_value(&order).clone();
2139        let end = require_some(
2140            observed.iter().position(|name| name == "agent_end"),
2141            "public agent_end for the failed run",
2142        )?;
2143        let settled = require_some(
2144            observed.iter().position(|name| name == "agent_settled"),
2145            "agent_settled after the failed run",
2146        )?;
2147        assert!(
2148            end < settled,
2149            "agent_end must be observed before agent_settled: {observed:?}"
2150        );
2151        Ok(())
2152    }
2153
2154    #[tokio::test]
2155    async fn idle_custom_message_append_failure_publishes_nothing() -> TestResult {
2156        let dir = tempfile::tempdir().test_context("tempdir")?;
2157        let provider = Arc::new(SeqProvider::new(one(start_event())));
2158        let session = blocked_append_session(provider, &dir)?;
2159        let events = Arc::new(StdMutex::new(Vec::new()));
2160        let events_clone = Arc::clone(&events);
2161        let _unsub = session.subscribe(move |event| {
2162            mutex_value(&events_clone).push(event.type_name().to_owned());
2163        });
2164        let transcript_before = session.messages().len();
2165
2166        let err = require_error(
2167            session
2168                .send_custom_message(
2169                    CustomMessageInput {
2170                        custom_type: "note".to_owned(),
2171                        content: CustomMessageContent::Text("hello".to_owned()),
2172                        display: true,
2173                        details: None,
2174                    },
2175                    false,
2176                    None,
2177                )
2178                .await,
2179            "idle custom append",
2180        )?;
2181        assert!(matches!(err, PromptError::Session(_)), "{err}");
2182        assert_eq!(
2183            session.messages().len(),
2184            transcript_before,
2185            "failed durable append must not mutate live transcript"
2186        );
2187        assert!(
2188            mutex_value(&events).is_empty(),
2189            "failed durable append must not publish message events: {:?}",
2190            mutex_value(&events)
2191        );
2192        Ok(())
2193    }
2194
2195    #[tokio::test]
2196    async fn message_end_disk_failure_settles_only_after_agent_end() -> TestResult {
2197        let dir = tempfile::tempdir().test_context("tempdir")?;
2198        let provider = Arc::new(SeqProvider::new(two(
2199            start_event(),
2200            done_ok(assistant_text("answer")),
2201        )));
2202        let session = blocked_append_session(provider, &dir)?;
2203        let order = Arc::new(StdMutex::new(Vec::new()));
2204        let order_clone = Arc::clone(&order);
2205        let _unsub = session.subscribe(move |event| {
2206            mutex_value(&order_clone).push(event.type_name().to_owned());
2207        });
2208
2209        let err = require_error(
2210            session.prompt("question", PromptOptions::default()).await,
2211            "disk-failure run",
2212        )?;
2213        assert!(matches!(err, PromptError::Session(_)), "{err}");
2214
2215        let observed = mutex_value(&order).clone();
2216        let message_end = require_some(
2217            observed.iter().position(|name| name == "message_end"),
2218            "public message_end for the failed persistence",
2219        )?;
2220        let agent_end = require_some(
2221            observed.iter().position(|name| name == "agent_end"),
2222            "public agent_end after the persistence failure",
2223        )?;
2224        let settled = require_some(
2225            observed.iter().position(|name| name == "agent_settled"),
2226            "final agent_settled",
2227        )?;
2228        assert!(
2229            message_end < agent_end && agent_end < settled,
2230            "persistence failure must not settle before agent_end: {observed:?}"
2231        );
2232        assert_eq!(
2233            observed
2234                .iter()
2235                .filter(|name| *name == "agent_settled")
2236                .count(),
2237            1,
2238            "exactly one settle per run: {observed:?}"
2239        );
2240        Ok(())
2241    }
2242
2243    #[allow(dead_code)]
2244    fn _ensure_null_runner_send_sync(_: NullExtensionRunner) {}
2245
2246    #[allow(dead_code)]
2247    fn _ensure_ok_event_ok(e: AssistantMessageEvent) {
2248        let _ = ok_event(e);
2249    }
2250
2251    #[tokio::test]
2252    async fn cancelled_preflight_releases_run_admission() -> TestResult {
2253        let provider = Arc::new(SeqProvider::new(two(
2254            start_event(),
2255            done_ok(assistant_text("after cancel")),
2256        )));
2257        let gate = Arc::new(Semaphore::new(0));
2258        let entered = Arc::new(Notify::new());
2259        let runner = Arc::new(TestRunner {
2260            before_start_gate: Some(Arc::clone(&gate)),
2261            before_start_entered: Some(Arc::clone(&entered)),
2262            ..TestRunner::default()
2263        });
2264        let mut config = AgentSessionConfig::test_config(provider.clone(), test_model())
2265            .test_context("test config")?;
2266        config.extension_runner = Some(runner);
2267        let session = AgentSession::new(config).test_context("session")?;
2268
2269        let entered_wait = entered.notified();
2270        let cancelled_session = Arc::clone(&session);
2271        let cancelled = tokio::spawn(async move {
2272            cancelled_session
2273                .prompt("cancelled", PromptOptions::default())
2274                .await
2275        });
2276        entered_wait.await;
2277        cancelled.abort();
2278        let join_error = require_error(cancelled.await, "cancelled prompt task")?;
2279        assert!(join_error.is_cancelled());
2280
2281        gate.add_permits(1);
2282        session
2283            .prompt("after cancel", PromptOptions::default())
2284            .await
2285            .test_context("prompt after cancellation")?;
2286        assert_eq!(provider.call_count(), 1);
2287        Ok(())
2288    }
2289}