1use 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
26const FAILED_RUN_END_BARRIER_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(1);
34
35#[derive(Clone, Copy, Debug, Eq, PartialEq)]
37pub enum StreamingBehavior {
38 Steer,
40 FollowUp,
42}
43
44impl StreamingBehavior {
45 #[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
55pub type PreflightCallback = Arc<dyn Fn(bool) + Send + Sync>;
61
62#[derive(Clone)]
64pub struct PromptOptions {
65 pub images: Vec<ImageContent>,
67 pub streaming_behavior: Option<StreamingBehavior>,
70 pub source: Option<String>,
72 pub expand_prompt_templates: bool,
74 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 #[must_use]
93 pub fn new() -> Self {
94 Self::default()
95 }
96}
97
98#[derive(Clone, Copy, Debug, Eq, PartialEq)]
100pub enum DeliverAs {
101 Steer,
103 FollowUp,
105 NextTurn,
107}
108
109#[derive(Clone, Debug)]
111pub struct CustomMessageInput {
112 pub custom_type: String,
114 pub content: CustomMessageContent,
116 pub display: bool,
118 pub details: Option<serde_json::Value>,
120}
121
122#[derive(Debug, thiserror::Error)]
124pub enum PromptError {
125 #[error("{0}")]
127 Message(String),
128 #[error(transparent)]
130 Agent(#[from] pi_agent::AgentLoopError),
131 #[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 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 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 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 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 {
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 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 async fn prompt_preflight(
350 self: &Arc<Self>,
351 text: &str,
352 options: &mut PromptOptions,
353 expand: bool,
354 ) -> Result<PreflightOutcome, PromptError> {
355 if expand && text.starts_with('/') && self.try_execute_extension_command(text).await? {
357 return Ok(PreflightOutcome::Handled);
358 }
359
360 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 let expanded_text = if expand {
372 self.expand_text(¤t_text)
373 } else {
374 current_text
375 };
376
377 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 self.flush_pending_bash_messages()
398 .await
399 .map_err(map_bash_flush_error)?;
400
401 let model = self.model();
403 if is_no_model(&model) {
404 return Err(PromptError::Message(format_no_model_selected_message()));
405 }
406
407 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 if let Some(last_msg) = self.agent.last_assistant() {
427 self.check_compaction(&last_msg, false).await;
428 }
429
430 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 let runner = self.hooks.runner();
442 let images_for_ext = if current_images.is_empty() {
443 None
444 } else {
445 serde_json::to_value(¤t_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 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 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 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 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 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 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 if self.check_compaction(&msg, true).await {
676 return Ok(true);
677 }
678
679 Ok(self.agent.has_queued_messages())
680 }
681
682 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 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 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
820fn 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
835fn 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 async fn drain(session: &Arc<AgentSession>) {
1040 session.wait_for_idle().await;
1041 }
1042
1043 fn ok_event(e: AssistantMessageEvent) -> AssistantMessageEvent {
1045 e
1046 }
1047
1048 fn sequence(events: Vec<AssistantMessageEvent>) -> ProviderResponse {
1050 events.into_iter().map(Ok).collect()
1051 }
1052
1053 fn one(e: AssistantMessageEvent) -> ProviderResponses {
1055 vec![sequence(vec![e])]
1056 }
1057
1058 fn two(a: AssistantMessageEvent, b: AssistantMessageEvent) -> ProviderResponses {
1060 vec![sequence(vec![a, b])]
1061 }
1062
1063 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 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 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 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 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 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 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}