1use crate::core::resources::{
26 ResourceLoader, SlashCommandInfo, SlashCommandSource, SyntheticSourceInfoOptions,
27 create_synthetic_source_info,
28};
29use std::sync::Arc;
30#[cfg(test)]
31use std::sync::Mutex;
32
33use super::AgentSession;
34use super::events::{
35 AgentSessionEvent, SessionShutdownReason, SessionStartEvent, SessionStartReason,
36};
37use super::prompt::{CustomMessageInput, DeliverAs, PromptError};
38use pi_ai::ImageContent;
39
40#[cfg(test)]
45pub(super) type ReloadRestartFactory = Arc<
46 dyn Fn(
47 Vec<String>,
48 String,
49 bool,
50 ) -> futures::future::BoxFuture<
51 'static,
52 Result<
53 Arc<crate::core::extension_host::HostExtensionRunner>,
54 crate::core::extension_host::HostStartError,
55 >,
56 > + Send
57 + Sync,
58>;
59
60#[derive(Clone, Copy, Debug, Eq, PartialEq)]
62pub enum ExtensionMode {
63 Tui,
65 Print,
67 Json,
69 Rpc,
71}
72
73impl ExtensionMode {
74 #[must_use]
76 pub const fn as_str(self) -> &'static str {
77 match self {
78 Self::Tui => "tui",
79 Self::Print => "print",
80 Self::Json => "json",
81 Self::Rpc => "rpc",
82 }
83 }
84}
85
86#[derive(Clone, Debug, Default)]
91pub struct ExtensionUiContext {
92 pub tag: Option<String>,
94}
95
96pub type ExtensionErrorListener = Arc<dyn Fn(&str, &str, &str) + Send + Sync>;
98
99#[derive(Clone, Default)]
101pub struct ExtensionBindings {
102 pub ui_context: Option<ExtensionUiContext>,
104 pub mode: Option<ExtensionMode>,
106 pub command_context_actions: Option<serde_json::Value>,
108 pub shutdown_handler: Option<Arc<dyn Fn() + Send + Sync>>,
110 pub on_error: Option<ExtensionErrorListener>,
112}
113
114#[derive(Clone)]
119pub struct ReplacedSessionContext {
120 pub session_id: String,
122 session: Arc<AgentSession>,
123}
124
125impl ReplacedSessionContext {
126 pub async fn send_custom_message(
132 &self,
133 message: CustomMessageInput,
134 trigger_turn: bool,
135 deliver_as: Option<DeliverAs>,
136 ) -> Result<(), PromptError> {
137 self.session
138 .send_custom_message(message, trigger_turn, deliver_as)
139 .await
140 }
141
142 pub async fn send_user_message(
148 &self,
149 text: &str,
150 images: Vec<ImageContent>,
151 deliver_as: Option<DeliverAs>,
152 ) -> Result<(), PromptError> {
153 self.session
154 .send_user_message(text, images, deliver_as)
155 .await
156 }
157}
158
159#[derive(Debug, thiserror::Error)]
161pub enum ExtensionBindError {
162 #[error("extension resource discovery failed: {0}")]
164 ResourceDiscover(super::extension_runner::ExtensionRunnerError),
165 #[error("resource reload failed: {0}")]
167 ResourceReload(String),
168 #[error("extension host restart failed: {0}")]
170 HostRestart(String),
171}
172
173impl AgentSession {
174 #[must_use]
178 pub fn has_extension_handlers(&self, event_type: &str) -> bool {
179 self.hooks.runner().has_handlers(event_type)
180 }
181
182 pub async fn bind_extensions(
196 &self,
197 bindings: ExtensionBindings,
198 ) -> Result<(), ExtensionBindError> {
199 let _bind_guard = self.bind_lock.lock().await;
200 {
202 let mut inner = self.lock_inner();
203 inner.extension_mode = bindings.mode;
204 inner.extension_ui_tag = bindings.ui_context.as_ref().and_then(|c| c.tag.clone());
205 inner
206 .extension_shutdown_handler
207 .clone_from(&bindings.shutdown_handler);
208 inner
209 .extension_error_listener
210 .clone_from(&bindings.on_error);
211 inner
212 .extension_command_context
213 .clone_from(&bindings.command_context_actions);
214 }
215
216 let pending = self.lock_inner().pending_session_start.take();
217 if let Some(event) = &pending {
218 let _ = self
221 .hooks
222 .runner()
223 .emit(AgentSessionEvent::SessionStart {
224 reason: event.reason,
225 previous_session_file: event.previous_session_file.clone(),
226 })
227 .await;
228 }
229 let discover_reason = match pending {
232 Some(SessionStartEvent {
233 reason: SessionStartReason::Reload,
234 ..
235 }) => "reload",
236 _ => "startup",
237 };
238 self.extend_resources_from_extensions(discover_reason)
239 .await?;
240 Ok(())
241 }
242
243 pub async fn extend_resources_from_extensions(
252 &self,
253 reason: &str,
254 ) -> Result<(), ExtensionBindError> {
255 if reason == SessionStartReason::Startup.as_str()
256 && self.lock_inner().initial_resources_discovered
257 {
258 return Ok(());
259 }
260
261 if reason == SessionStartReason::Reload.as_str()
262 && let Some(loader) = &self.resource_loader
263 {
264 let mut loader = loader.lock().await;
265 loader
266 .reload()
267 .await
268 .map_err(|error| ExtensionBindError::ResourceReload(error.to_string()))?;
269 self.apply_resource_snapshot(&loader);
270 }
271
272 let runner = self.hooks.runner();
273 if !runner.has_handlers("resources_discover") {
274 return Ok(());
275 }
276 let paths = runner
277 .emit_resources_discover(&self.cwd, reason)
278 .await
279 .map_err(ExtensionBindError::ResourceDiscover)?;
280 if let Some(loader) = &self.resource_loader {
281 let mut loader = loader.lock().await;
282 loader.extend_resources(paths);
283 self.apply_resource_snapshot(&loader);
284 }
285 if reason == SessionStartReason::Startup.as_str() {
286 self.lock_inner().initial_resources_discovered = true;
287 }
288 Ok(())
289 }
290
291 fn apply_resource_snapshot(&self, loader: &crate::core::resources::DefaultResourceLoader) {
292 let skills = loader.get_skills().0.to_vec();
293 let prompt_templates = loader.get_prompts().0.to_vec();
294 let append = (!loader.get_append_system_prompt().is_empty())
295 .then(|| loader.get_append_system_prompt().join("\n\n"));
296 let selected_tools = self.lock_inner().active_tool_names.clone();
297 let system_prompt = crate::core::system_prompt::build_system_prompt(
298 &crate::core::system_prompt::BuildSystemPromptOptions {
299 custom_prompt: loader.get_system_prompt().map(str::to_owned),
300 selected_tools: Some(selected_tools),
301 append,
302 cwd: self.cwd.clone(),
303 context_files: Some(loader.get_agents_files().to_vec()),
304 skills: Some(skills.clone()),
305 ..crate::core::system_prompt::BuildSystemPromptOptions::default()
306 },
307 );
308 *self
309 .skills
310 .lock()
311 .unwrap_or_else(std::sync::PoisonError::into_inner) = skills;
312 *self
313 .prompt_templates
314 .lock()
315 .unwrap_or_else(std::sync::PoisonError::into_inner) = prompt_templates;
316 self.lock_inner()
317 .base_system_prompt
318 .clone_from(&system_prompt);
319 self.hooks.set_base_system_prompt(system_prompt.clone());
320 self.hooks.set_system_prompt_override(None);
321 self.agent.set_system_prompt(system_prompt);
322 }
323
324 #[must_use]
329 pub fn get_extension_source_label(extension_path: &str) -> String {
330 if extension_path.starts_with('<') {
331 let trimmed = extension_path.trim_start_matches('<').trim_end_matches('>');
332 return format!("extension:{trimmed}");
333 }
334 let base = std::path::Path::new(extension_path)
335 .file_stem()
336 .map_or_else(
337 || extension_path.to_owned(),
338 |s| s.to_string_lossy().into_owned(),
339 );
340 format!("extension:{base}")
341 }
342
343 pub async fn reload(&self) -> Result<(), ExtensionBindError> {
360 let runner = self.hooks.runner();
361 let previous_flag_values = runner.get_flag_values();
362
363 let _ = runner
366 .emit(AgentSessionEvent::SessionShutdown {
367 reason: SessionShutdownReason::Reload,
368 target_session_file: None,
369 })
370 .await;
371
372 if let Some(host) = self.host_extension_runner() {
373 let Some(runtime) = self.model_runtime() else {
374 self.hooks
378 .set_runner(Arc::new(super::extension_runner::NullExtensionRunner));
379 self.set_host_extension_runner(None);
380 self.refresh_tool_registry(&super::tools::RefreshToolRegistryOptions {
381 active_tool_names: None,
382 include_all_extension_tools: true,
383 });
384 host.retire_after_cutover().await;
385 self.emit_session_start_reload().await;
386 self.extend_resources_from_extensions(SessionStartReason::Reload.as_str())
387 .await?;
388 return Ok(());
389 };
390 let new_host = match self
391 .restart_host_for_reload(&host, &runtime, previous_flag_values)
392 .await
393 {
394 Ok(host) => host,
395 Err(error) => {
396 self.emit_session_start_reload().await;
400 return Err(ExtensionBindError::HostRestart(error.to_string()));
401 }
402 };
403 self.hooks.set_runner(
409 Arc::clone(&new_host) as Arc<dyn super::extension_runner::ExtensionRunner>
410 );
411 self.set_host_extension_runner(Some(new_host));
412 self.refresh_tool_registry(&super::tools::RefreshToolRegistryOptions {
413 active_tool_names: None,
414 include_all_extension_tools: true,
415 });
416 host.retire_after_cutover().await;
419 self.emit_session_start_reload().await;
420 self.extend_resources_from_extensions(SessionStartReason::Reload.as_str())
421 .await?;
422 return Ok(());
423 }
424
425 self.emit_session_start_reload().await;
427 self.extend_resources_from_extensions(SessionStartReason::Reload.as_str())
428 .await?;
429 Ok(())
430 }
431
432 async fn restart_host_for_reload(
434 &self,
435 host: &Arc<crate::core::extension_host::HostExtensionRunner>,
436 runtime: &crate::core::model_runtime::ModelRuntime,
437 previous_flag_values: std::collections::HashMap<String, serde_json::Value>,
438 ) -> Result<
439 Arc<crate::core::extension_host::HostExtensionRunner>,
440 crate::core::extension_host::HostStartError,
441 > {
442 #[cfg(test)]
443 {
444 let factory = self
445 .lock_inner()
446 .reload_restart_factory
447 .as_ref()
448 .map(Arc::clone);
449 if let Some(factory) = factory {
450 return host
451 .restart_and_rewire_with(
452 runtime,
453 previous_flag_values,
454 move |paths, cwd, trusted| factory(paths, cwd, trusted),
455 )
456 .await;
457 }
458 }
459 host.restart_and_rewire(runtime, previous_flag_values).await
460 }
461
462 async fn emit_session_start_reload(&self) {
464 let _ = self
465 .hooks
466 .runner()
467 .emit(AgentSessionEvent::SessionStart {
468 reason: SessionStartReason::Reload,
469 previous_session_file: None,
470 })
471 .await;
472 }
473
474 pub async fn create_replaced_session_context(self: &Arc<Self>) -> ReplacedSessionContext {
479 let session_id = self.session_id().await;
480 ReplacedSessionContext {
481 session_id,
482 session: Arc::clone(self),
483 }
484 }
485
486 #[must_use]
488 pub fn slash_commands(&self) -> Vec<SlashCommandInfo> {
489 let mut commands = Vec::new();
490 let mut extension_names = std::collections::HashSet::new();
491
492 if let Some(host) = self.host_extension_runner() {
493 for command in host.registry().commands() {
494 extension_names.insert(command.name.clone());
495 let path = command
496 .source
497 .clone()
498 .unwrap_or_else(|| "<extension>".to_owned());
499 commands.push(SlashCommandInfo {
500 name: command.name.clone(),
501 description: command.description.clone(),
502 source: SlashCommandSource::Extension,
503 source_info: create_synthetic_source_info(
504 path,
505 SyntheticSourceInfoOptions {
506 source: "extension".to_owned(),
507 scope: None,
508 origin: None,
509 base_dir: None,
510 },
511 ),
512 });
513 }
514 }
515
516 for name in self.hooks.runner().get_registered_commands() {
517 if extension_names.insert(name.clone()) {
518 commands.push(SlashCommandInfo {
519 name,
520 description: None,
521 source: SlashCommandSource::Extension,
522 source_info: create_synthetic_source_info(
523 "<extension>",
524 SyntheticSourceInfoOptions {
525 source: "extension".to_owned(),
526 scope: None,
527 origin: None,
528 base_dir: None,
529 },
530 ),
531 });
532 }
533 }
534
535 commands.extend(
536 self.prompt_templates
537 .lock()
538 .unwrap_or_else(std::sync::PoisonError::into_inner)
539 .iter()
540 .map(|template| SlashCommandInfo {
541 name: template.name.clone(),
542 description: Some(template.description.clone()),
543 source: SlashCommandSource::Prompt,
544 source_info: template.source_info.clone(),
545 }),
546 );
547 commands.extend(
548 self.skills
549 .lock()
550 .unwrap_or_else(std::sync::PoisonError::into_inner)
551 .iter()
552 .map(|skill| SlashCommandInfo {
553 name: format!("skill:{}", skill.name),
554 description: Some(skill.description.clone()),
555 source: SlashCommandSource::Skill,
556 source_info: skill.source_info.clone(),
557 }),
558 );
559 commands
560 }
561
562 pub fn report_extension_error(&self, extension_path: &str, event: &str, error: &str) {
564 let listener = self.lock_inner().extension_error_listener.clone();
565 if let Some(listener) = listener {
566 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
567 listener(extension_path, event, error);
568 }));
569 }
570 }
571
572 pub fn invoke_extension_shutdown_handler(&self) {
574 let handler = self.lock_inner().extension_shutdown_handler.clone();
575 if let Some(handler) = handler {
576 handler();
577 }
578 }
579
580 #[must_use]
582 pub fn extension_mode(&self) -> Option<ExtensionMode> {
583 self.lock_inner().extension_mode
584 }
585
586 #[must_use]
588 pub fn host_extension_runner(
589 &self,
590 ) -> Option<Arc<crate::core::extension_host::HostExtensionRunner>> {
591 self.host_extension_runner
592 .read()
593 .ok()
594 .and_then(|guard| guard.clone())
595 }
596
597 pub fn set_host_extension_runner(
599 &self,
600 runner: Option<Arc<crate::core::extension_host::HostExtensionRunner>>,
601 ) {
602 if let Ok(mut guard) = self.host_extension_runner.write() {
603 *guard = runner;
604 }
605 }
606
607 #[cfg(test)]
609 pub(super) fn set_reload_restart_factory(&self, factory: Option<ReloadRestartFactory>) {
610 self.lock_inner().reload_restart_factory = factory;
611 }
612}
613
614#[cfg(test)]
619mod tests {
620 use super::*;
621 use crate::core::agent_session::extension_runner::ExtensionRunner;
622 use crate::core::agent_session::{AgentSession, AgentSessionConfig};
623 use futures::future::BoxFuture;
624 use futures::stream::{self, BoxStream, StreamExt};
625 use pi_ai::{
626 AssistantMessageEvent, Context, Model, ModelCost, ModelInput, Provider, ProviderError,
627 StreamOptions,
628 };
629 use std::collections::HashMap;
630 use std::collections::HashSet;
631 use std::error::Error;
632 use std::io;
633 use std::sync::Mutex as StdMutex;
634 use std::sync::atomic::{AtomicBool, Ordering};
635 use std::time::Duration;
636
637 use pi_ext::client::HostClient;
638 use pi_ext::protocol::{Frame, FrameKind, HelloAck, decode_frame_str, encode_frame};
639 use serde_json::{Map, Value, json};
640 use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
641 use tokio::sync::mpsc;
642
643 use crate::core::extension_host::{HostExtensionRunner, HostStartError};
644 use crate::core::model_runtime::ModelRuntime;
645
646 type TestResult<T = ()> = Result<T, Box<dyn Error>>;
647
648 fn test_model() -> Model {
649 Model {
650 id: "m".to_owned(),
651 name: "m".to_owned(),
652 api: "test-api".to_owned(),
653 provider: "test-provider".to_owned(),
654 base_url: String::new(),
655 reasoning: false,
656 thinking_level_map: None,
657 input: vec![ModelInput::Text],
658 cost: ModelCost::default(),
659 context_window: 8_192,
660 max_tokens: 1_024,
661 headers: None,
662 compat: None,
663 extra: std::collections::BTreeMap::new(),
664 }
665 }
666
667 #[derive(Clone)]
668 struct StubProvider;
669
670 impl Provider for StubProvider {
671 fn stream(
672 &self,
673 _model: &Model,
674 _context: Context,
675 _options: StreamOptions,
676 ) -> BoxStream<'static, Result<AssistantMessageEvent, ProviderError>> {
677 stream::empty().boxed()
678 }
679 }
680
681 fn make_session() -> Result<Arc<AgentSession>, crate::core::agent_session::AgentSessionError> {
682 let config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
683 AgentSession::new(config)
684 }
685
686 fn locked_clone<T: Clone>(value: &Mutex<T>, label: &str) -> TestResult<T> {
687 value
688 .lock()
689 .map(|guard| guard.clone())
690 .map_err(|_| io::Error::other(format!("{label} mutex poisoned")).into())
691 }
692
693 struct TestRunner {
696 has_start: AtomicBool,
697 has_shutdown: AtomicBool,
698 has_resources: AtomicBool,
699 calls: Arc<Mutex<Vec<String>>>,
703 emit_delay: Mutex<Option<std::time::Duration>>,
704 flag_values: Arc<Mutex<HashMap<String, serde_json::Value>>>,
705 resource_paths: Arc<Mutex<crate::core::resources::ResourceExtensionPaths>>,
706 }
707
708 impl TestRunner {
709 fn new() -> Self {
710 Self {
711 has_start: AtomicBool::new(false),
712 has_shutdown: AtomicBool::new(false),
713 has_resources: AtomicBool::new(false),
714 calls: Arc::new(Mutex::new(Vec::new())),
715 emit_delay: Mutex::new(None),
716 flag_values: Arc::new(Mutex::new(HashMap::new())),
717 resource_paths: Arc::new(Mutex::new(
718 crate::core::resources::ResourceExtensionPaths::default(),
719 )),
720 }
721 }
722
723 fn record(&self, entry: String) {
724 if let Ok(mut g) = self.calls.lock() {
725 g.push(entry);
726 }
727 }
728
729 fn lifecycle_label(event: &AgentSessionEvent) -> String {
730 match event {
731 AgentSessionEvent::SessionStart {
732 reason,
733 previous_session_file,
734 } => format!(
735 "session_start:{}:{}",
736 reason.as_str(),
737 previous_session_file.as_deref().unwrap_or("-")
738 ),
739 AgentSessionEvent::SessionShutdown {
740 reason,
741 target_session_file,
742 } => format!(
743 "session_shutdown:{}:{}",
744 reason.as_str(),
745 target_session_file.as_deref().unwrap_or("-")
746 ),
747 other => other.type_name().to_owned(),
748 }
749 }
750 }
751
752 impl ExtensionRunner for TestRunner {
753 fn has_handlers(&self, event: &str) -> bool {
754 match event {
755 "session_start" => self.has_start.load(Ordering::SeqCst),
756 "session_shutdown" => self.has_shutdown.load(Ordering::SeqCst),
757 "resources_discover" => self.has_resources.load(Ordering::SeqCst),
758 _ => false,
759 }
760 }
761
762 fn emit(
763 &self,
764 event: AgentSessionEvent,
765 ) -> BoxFuture<
766 '_,
767 Result<
768 Option<super::super::extension_runner::CancelResult>,
769 super::super::extension_runner::ExtensionRunnerError,
770 >,
771 > {
772 let delay = self
773 .emit_delay
774 .lock()
775 .map(|guard| *guard)
776 .unwrap_or_default();
777 let label = Self::lifecycle_label(&event);
778 Box::pin(async move {
779 if let Some(delay) = delay {
780 tokio::time::sleep(delay).await;
781 }
782 self.record(label);
783 Ok(None)
784 })
785 }
786
787 fn emit_message_end(
788 &self,
789 message: pi_agent::AgentMessage,
790 ) -> BoxFuture<
791 '_,
792 Result<
793 Option<pi_agent::AgentMessage>,
794 super::super::extension_runner::ExtensionRunnerError,
795 >,
796 > {
797 Box::pin(async move { Ok(Some(message)) })
798 }
799
800 fn emit_tool_call(
801 &self,
802 _tool_name: &str,
803 _tool_call_id: &str,
804 _input: serde_json::Map<String, serde_json::Value>,
805 ) -> BoxFuture<
806 '_,
807 Result<
808 Option<pi_agent::BeforeToolCallResult>,
809 super::super::extension_runner::ExtensionRunnerError,
810 >,
811 > {
812 Box::pin(async { Ok(None) })
813 }
814
815 fn emit_tool_result(
816 &self,
817 _tool_name: &str,
818 _tool_call_id: &str,
819 _input: serde_json::Map<String, serde_json::Value>,
820 _content: Vec<pi_ai::ToolResultContent>,
821 _details: serde_json::Value,
822 _is_error: bool,
823 ) -> BoxFuture<
824 '_,
825 Result<
826 Option<pi_agent::AfterToolCallResult>,
827 super::super::extension_runner::ExtensionRunnerError,
828 >,
829 > {
830 Box::pin(async { Ok(None) })
831 }
832
833 fn emit_input(
834 &self,
835 _text: &str,
836 _images: Option<serde_json::Value>,
837 _source: &str,
838 _streaming_behavior: Option<&str>,
839 ) -> BoxFuture<
840 '_,
841 Result<
842 super::super::extension_runner::InputTransformResult,
843 super::super::extension_runner::ExtensionRunnerError,
844 >,
845 > {
846 Box::pin(async { Ok(super::super::extension_runner::InputTransformResult::default()) })
847 }
848
849 fn emit_before_agent_start(
850 &self,
851 _prompt: &str,
852 _images: Option<serde_json::Value>,
853 ) -> BoxFuture<
854 '_,
855 Result<
856 Option<super::super::extension_runner::BeforeAgentStartResult>,
857 super::super::extension_runner::ExtensionRunnerError,
858 >,
859 > {
860 Box::pin(async { Ok(None) })
861 }
862
863 fn emit_resources_discover(
864 &self,
865 cwd: &str,
866 reason: &str,
867 ) -> BoxFuture<
868 '_,
869 Result<
870 crate::core::resources::ResourceExtensionPaths,
871 super::super::extension_runner::ExtensionRunnerError,
872 >,
873 > {
874 self.record(format!("resources_discover:{reason}"));
875 let _ = cwd;
876 let paths = self
877 .resource_paths
878 .lock()
879 .map(|paths| paths.clone())
880 .unwrap_or_default();
881 Box::pin(async move { Ok(paths) })
882 }
883
884 fn get_registered_commands(&self) -> Vec<String> {
885 Vec::new()
886 }
887
888 fn get_all_registered_tools(&self) -> HashMap<String, Arc<dyn pi_agent::AgentTool>> {
889 HashMap::new()
890 }
891
892 fn get_flag_values(&self) -> HashMap<String, serde_json::Value> {
893 self.flag_values
894 .lock()
895 .map(|g| g.clone())
896 .unwrap_or_default()
897 }
898 fn execute_command(
899 &self,
900 _name: &str,
901 _args: &str,
902 ) -> BoxFuture<'_, Result<bool, super::super::extension_runner::ExtensionRunnerError>>
903 {
904 Box::pin(async { Ok(false) })
905 }
906
907 fn invalidate(&self) {}
908
909 fn emit_error(&self, _message: String) {}
910 }
911
912 #[tokio::test]
913 async fn has_extension_handlers_delegates_to_runner() -> TestResult {
914 let runner = Arc::new(TestRunner::new());
915 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
916 config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
917 let session = AgentSession::new(config)?;
918 assert!(!session.has_extension_handlers("session_start"));
919 runner.has_start.store(true, Ordering::SeqCst);
920 assert!(session.has_extension_handlers("session_start"));
921 Ok(())
922 }
923
924 #[tokio::test]
925 async fn bind_extensions_records_bindings_and_discovers_resources() -> TestResult {
926 let runner = Arc::new(TestRunner::new());
927 runner.has_resources.store(true, Ordering::SeqCst);
928 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
929 config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
930 let session = AgentSession::new(config)?;
931
932 let error_hit = Arc::new(AtomicBool::new(false));
933 let error_hit_clone = Arc::clone(&error_hit);
934 let bindings = ExtensionBindings {
935 mode: Some(ExtensionMode::Rpc),
936 on_error: Some(Arc::new(move |_path: &str, _event: &str, _error: &str| {
937 error_hit_clone.store(true, Ordering::SeqCst);
938 })),
939 ..Default::default()
940 };
941 session.bind_extensions(bindings).await?;
942
943 let calls = locked_clone(&runner.calls, "calls")?;
945 assert!(calls.iter().any(|c| c == "resources_discover:startup"));
946
947 assert_eq!(session.extension_mode(), Some(ExtensionMode::Rpc));
949
950 session.report_extension_error("extension.ts", "agent_start", "boom");
952 assert!(error_hit.load(Ordering::SeqCst));
953 Ok(())
954 }
955
956 #[tokio::test]
957 async fn bind_emits_stored_session_start_before_discovery() -> TestResult {
958 let runner = Arc::new(TestRunner::new());
959 runner.has_start.store(true, Ordering::SeqCst);
960 runner.has_resources.store(true, Ordering::SeqCst);
961 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
962 config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
963 let session = AgentSession::new(config)?;
964
965 session
966 .bind_extensions(ExtensionBindings::default())
967 .await?;
968 let calls = locked_clone(&runner.calls, "calls")?;
969 assert_eq!(
970 calls,
971 vec![
972 "session_start:startup:-".to_owned(),
973 "resources_discover:startup".to_owned(),
974 ],
975 "first bind must emit session_start exactly once, before discovery"
976 );
977
978 session
980 .bind_extensions(ExtensionBindings::default())
981 .await?;
982 let calls = locked_clone(&runner.calls, "calls")?;
983 assert_eq!(
984 calls,
985 vec![
986 "session_start:startup:-".to_owned(),
987 "resources_discover:startup".to_owned(),
988 ],
989 "second bind must not re-emit or rediscover"
990 );
991 Ok(())
992 }
993
994 #[tokio::test]
995 async fn concurrent_binds_emit_start_once_before_any_discovery() -> TestResult {
996 let runner = Arc::new(TestRunner::new());
997 runner.has_start.store(true, Ordering::SeqCst);
998 runner.has_resources.store(true, Ordering::SeqCst);
999 *runner
1000 .emit_delay
1001 .lock()
1002 .map_err(|_| io::Error::other("emit delay mutex poisoned"))? =
1003 Some(std::time::Duration::from_millis(25));
1004 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1005 config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
1006 let session = AgentSession::new(config)?;
1007
1008 let (first, second) = tokio::join!(
1009 session.bind_extensions(ExtensionBindings::default()),
1010 session.bind_extensions(ExtensionBindings::default()),
1011 );
1012 first?;
1013 second?;
1014
1015 let calls = locked_clone(&runner.calls, "calls")?;
1016 assert_eq!(
1017 calls,
1018 vec![
1019 "session_start:startup:-".to_owned(),
1020 "resources_discover:startup".to_owned(),
1021 ],
1022 "concurrent binds must serialize: one start strictly before one discovery"
1023 );
1024 Ok(())
1025 }
1026
1027 #[tokio::test]
1028 async fn bind_emits_replacement_reason_and_previous_file() -> TestResult {
1029 let runner = Arc::new(TestRunner::new());
1030 runner.has_start.store(true, Ordering::SeqCst);
1031 runner.has_resources.store(true, Ordering::SeqCst);
1032 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1033 config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
1034 config.session_start_event = Some(SessionStartEvent {
1035 reason: SessionStartReason::New,
1036 previous_session_file: Some("prev.jsonl".into()),
1037 });
1038 let session = AgentSession::new(config)?;
1039
1040 session
1041 .bind_extensions(ExtensionBindings::default())
1042 .await?;
1043 let calls = locked_clone(&runner.calls, "calls")?;
1044 assert_eq!(
1045 calls,
1046 vec![
1047 "session_start:new:prev.jsonl".to_owned(),
1048 "resources_discover:startup".to_owned(),
1051 ]
1052 );
1053 Ok(())
1054 }
1055
1056 #[tokio::test]
1057 async fn extension_resources_refresh_session_skills_and_system_prompt() -> TestResult {
1058 let temp = tempfile::tempdir()?;
1059 let cwd = temp.path().join("project");
1060 let agent_dir = temp.path().join("agent");
1061 let extension_dir = temp.path().join("extension");
1062 let skill_dir = extension_dir.join("skills");
1063 std::fs::create_dir_all(&cwd)?;
1064 std::fs::create_dir_all(&agent_dir)?;
1065 std::fs::create_dir_all(&skill_dir)?;
1066 std::fs::write(
1067 skill_dir.join("SKILL.md"),
1068 "---\nname: extension-skill\ndescription: extension\n---\nbody\n",
1069 )?;
1070 let extension_path = extension_dir.join("plugin.ts");
1071 std::fs::write(&extension_path, "")?;
1072 let extension_path = extension_path.to_string_lossy().into_owned();
1073
1074 let runner = Arc::new(TestRunner::new());
1075 runner.has_resources.store(true, Ordering::SeqCst);
1076 *runner
1077 .resource_paths
1078 .lock()
1079 .map_err(|_| io::Error::other("resource paths mutex poisoned"))? =
1080 crate::core::resources::ResourceExtensionPaths {
1081 skill_paths: vec![crate::core::resources::ExtensionResourcePath::discovered(
1082 "skills".to_owned(),
1083 &extension_path,
1084 )],
1085 ..crate::core::resources::ResourceExtensionPaths::default()
1086 };
1087
1088 let settings = crate::core::settings::SettingsManager::create(
1089 &cwd,
1090 Some(&agent_dir),
1091 crate::core::settings::SettingsManagerCreateOptions::new().project_trusted(true),
1092 );
1093 let mut loader = crate::core::resources::DefaultResourceLoader::new(
1094 crate::core::resources::DefaultResourceLoaderOptions {
1095 cwd: cwd.clone(),
1096 agent_dir,
1097 settings_manager: Some(settings),
1098 ..crate::core::resources::DefaultResourceLoaderOptions::default()
1099 },
1100 );
1101 loader.reload().await?;
1102 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1103 config.cwd = cwd.to_string_lossy().into_owned();
1104 config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
1105 config.initial_active_tool_names = Some(vec!["read".to_owned()]);
1106 config.tools = vec![Arc::new(crate::core::tools::read::ReadTool::new(&cwd))];
1107 config.resource_loader = Some(loader);
1108 config.system_prompt = "stale".to_owned();
1109 let session = AgentSession::new(config)?;
1110
1111 session
1112 .bind_extensions(ExtensionBindings::default())
1113 .await?;
1114 assert!(
1115 session
1116 .agent
1117 .state()
1118 .system_prompt
1119 .contains("<name>extension-skill</name>")
1120 );
1121 assert!(
1122 session
1123 .skills
1124 .lock()
1125 .map_err(|_| io::Error::other("skills mutex poisoned"))?
1126 .iter()
1127 .any(|skill| skill.name == "extension-skill")
1128 );
1129
1130 *runner
1131 .resource_paths
1132 .lock()
1133 .map_err(|_| io::Error::other("resource paths mutex poisoned"))? =
1134 crate::core::resources::ResourceExtensionPaths::default();
1135 session.reload().await?;
1136 assert!(
1137 !session
1138 .agent
1139 .state()
1140 .system_prompt
1141 .contains("extension-skill")
1142 );
1143 assert!(
1144 session
1145 .skills
1146 .lock()
1147 .map_err(|_| io::Error::other("skills mutex poisoned"))?
1148 .is_empty()
1149 );
1150 Ok(())
1151 }
1152
1153 #[tokio::test]
1154 async fn reload_emits_shutdown_start_discovery_in_order() -> TestResult {
1155 let runner = Arc::new(TestRunner::new());
1156 runner.has_shutdown.store(true, Ordering::SeqCst);
1157 runner.has_start.store(true, Ordering::SeqCst);
1158 runner.has_resources.store(true, Ordering::SeqCst);
1159 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1160 config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
1161 let session = AgentSession::new(config)?;
1162
1163 session
1164 .bind_extensions(ExtensionBindings {
1165 mode: Some(ExtensionMode::Rpc),
1166 ..Default::default()
1167 })
1168 .await?;
1169 runner
1170 .calls
1171 .lock()
1172 .map_err(|_| io::Error::other("calls mutex poisoned"))?
1173 .clear();
1174
1175 session.reload().await?;
1176
1177 let calls = locked_clone(&runner.calls, "calls")?;
1178 assert_eq!(
1179 calls,
1180 vec![
1181 "session_shutdown:reload:-".to_owned(),
1182 "session_start:reload:-".to_owned(),
1183 "resources_discover:reload".to_owned(),
1184 ],
1185 "reload must emit shutdown, then start, then rediscover"
1186 );
1187 Ok(())
1188 }
1189
1190 #[tokio::test]
1191 async fn reload_rediscovers_resources_without_bindings() -> TestResult {
1192 let runner = Arc::new(TestRunner::new());
1193 runner.has_shutdown.store(true, Ordering::SeqCst);
1194 runner.has_resources.store(true, Ordering::SeqCst);
1195 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1196 config.extension_runner = Some(runner.clone() as Arc<dyn ExtensionRunner>);
1197 let session = AgentSession::new(config)?;
1198 session.reload().await?;
1199 let calls = locked_clone(&runner.calls, "calls")?;
1200 assert!(calls.iter().any(|c| c == "session_shutdown:reload:-"));
1201 assert!(
1202 calls.iter().any(|c| c == "resources_discover:reload"),
1203 "reload must refresh resources even without bindings, got {calls:?}"
1204 );
1205 Ok(())
1206 }
1207
1208 struct FailingRunner;
1210
1211 impl ExtensionRunner for FailingRunner {
1212 fn has_handlers(&self, _event: &str) -> bool {
1213 true
1214 }
1215
1216 fn emit(
1219 &self,
1220 _event: AgentSessionEvent,
1221 ) -> BoxFuture<
1222 '_,
1223 Result<
1224 Option<super::super::extension_runner::CancelResult>,
1225 super::super::extension_runner::ExtensionRunnerError,
1226 >,
1227 > {
1228 Box::pin(async {
1229 Err(
1230 super::super::extension_runner::ExtensionRunnerError::Failed(
1231 "host gone".into(),
1232 ),
1233 )
1234 })
1235 }
1236
1237 fn emit_message_end(
1238 &self,
1239 message: pi_agent::AgentMessage,
1240 ) -> BoxFuture<
1241 '_,
1242 Result<
1243 Option<pi_agent::AgentMessage>,
1244 super::super::extension_runner::ExtensionRunnerError,
1245 >,
1246 > {
1247 Box::pin(async move { Ok(Some(message)) })
1248 }
1249
1250 fn emit_tool_call(
1251 &self,
1252 _tool_name: &str,
1253 _tool_call_id: &str,
1254 _input: serde_json::Map<String, serde_json::Value>,
1255 ) -> BoxFuture<
1256 '_,
1257 Result<
1258 Option<pi_agent::BeforeToolCallResult>,
1259 super::super::extension_runner::ExtensionRunnerError,
1260 >,
1261 > {
1262 Box::pin(async { Ok(None) })
1263 }
1264
1265 fn emit_tool_result(
1266 &self,
1267 _tool_name: &str,
1268 _tool_call_id: &str,
1269 _input: serde_json::Map<String, serde_json::Value>,
1270 _content: Vec<pi_ai::ToolResultContent>,
1271 _details: serde_json::Value,
1272 _is_error: bool,
1273 ) -> BoxFuture<
1274 '_,
1275 Result<
1276 Option<pi_agent::AfterToolCallResult>,
1277 super::super::extension_runner::ExtensionRunnerError,
1278 >,
1279 > {
1280 Box::pin(async { Ok(None) })
1281 }
1282
1283 fn emit_input(
1284 &self,
1285 _text: &str,
1286 _images: Option<serde_json::Value>,
1287 _source: &str,
1288 _streaming_behavior: Option<&str>,
1289 ) -> BoxFuture<
1290 '_,
1291 Result<
1292 super::super::extension_runner::InputTransformResult,
1293 super::super::extension_runner::ExtensionRunnerError,
1294 >,
1295 > {
1296 Box::pin(async { Ok(super::super::extension_runner::InputTransformResult::default()) })
1297 }
1298
1299 fn emit_before_agent_start(
1300 &self,
1301 _prompt: &str,
1302 _images: Option<serde_json::Value>,
1303 ) -> BoxFuture<
1304 '_,
1305 Result<
1306 Option<super::super::extension_runner::BeforeAgentStartResult>,
1307 super::super::extension_runner::ExtensionRunnerError,
1308 >,
1309 > {
1310 Box::pin(async { Ok(None) })
1311 }
1312
1313 fn emit_resources_discover(
1314 &self,
1315 _cwd: &str,
1316 _reason: &str,
1317 ) -> BoxFuture<
1318 '_,
1319 Result<
1320 crate::core::resources::ResourceExtensionPaths,
1321 super::super::extension_runner::ExtensionRunnerError,
1322 >,
1323 > {
1324 Box::pin(async { Ok(crate::core::resources::ResourceExtensionPaths::default()) })
1325 }
1326
1327 fn get_registered_commands(&self) -> Vec<String> {
1328 Vec::new()
1329 }
1330
1331 fn get_all_registered_tools(&self) -> HashMap<String, Arc<dyn pi_agent::AgentTool>> {
1332 HashMap::new()
1333 }
1334
1335 fn get_flag_values(&self) -> HashMap<String, serde_json::Value> {
1336 HashMap::new()
1337 }
1338
1339 fn execute_command(
1340 &self,
1341 _name: &str,
1342 _args: &str,
1343 ) -> BoxFuture<'_, Result<bool, super::super::extension_runner::ExtensionRunnerError>>
1344 {
1345 Box::pin(async { Ok(false) })
1346 }
1347
1348 fn invalidate(&self) {}
1349 fn emit_error(&self, _message: String) {}
1350 }
1351
1352 #[tokio::test]
1353 async fn reload_survives_lifecycle_emit_error() -> TestResult {
1354 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1357 config.extension_runner = Some(Arc::new(FailingRunner) as Arc<dyn ExtensionRunner>);
1358 let session = AgentSession::new(config)?;
1359 session.reload().await?;
1360 Ok(())
1361 }
1362
1363 #[tokio::test]
1364 async fn bind_survives_lifecycle_emit_error() -> TestResult {
1365 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1366 config.extension_runner = Some(Arc::new(FailingRunner) as Arc<dyn ExtensionRunner>);
1367 let session = AgentSession::new(config)?;
1368 session
1369 .bind_extensions(ExtensionBindings::default())
1370 .await?;
1371 Ok(())
1372 }
1373
1374 #[tokio::test]
1375 async fn create_replaced_session_context_forwards_sends_to_session() -> TestResult {
1376 let session = make_session()?;
1377 let ctx = session.create_replaced_session_context().await;
1378 assert_eq!(ctx.session_id, session.session_id().await);
1379
1380 ctx.send_custom_message(
1381 crate::core::agent_session::prompt::CustomMessageInput {
1382 custom_type: "replacement-note".to_owned(),
1383 content: crate::core::messages::CustomMessageContent::Text("hi".to_owned()),
1384 display: true,
1385 details: None,
1386 },
1387 false,
1388 None,
1389 )
1390 .await?;
1391 ctx.send_user_message("yo", Vec::new(), None).await?;
1392
1393 let messages = session.agent.state().messages;
1394 assert!(messages.iter().any(|message| {
1395 matches!(
1396 message,
1397 pi_agent::AgentMessage::Custom(custom)
1398 if custom.payload.get("customType").and_then(serde_json::Value::as_str)
1399 == Some("replacement-note")
1400 )
1401 }));
1402 assert!(messages.iter().any(|message| {
1403 matches!(
1404 message,
1405 pi_agent::AgentMessage::Llm(message)
1406 if matches!(
1407 message.as_ref(),
1408 pi_ai::Message::User(user)
1409 if user.content == pi_ai::UserMessageContent::Text("yo".to_owned())
1410 )
1411 )
1412 }));
1413 Ok(())
1414 }
1415
1416 #[tokio::test]
1417 async fn get_extension_source_label_strips_extension() -> TestResult {
1418 assert_eq!(
1419 AgentSession::get_extension_source_label("/foo/bar/myext.ts"),
1420 "extension:myext"
1421 );
1422 assert_eq!(
1423 AgentSession::get_extension_source_label("<inline>"),
1424 "extension:inline"
1425 );
1426 Ok(())
1427 }
1428
1429 #[tokio::test]
1430 async fn extension_error_listener_receives_structured_fields_and_is_isolated() -> TestResult {
1431 let session = make_session()?;
1432 let received = Arc::new(Mutex::new(Vec::new()));
1433 let received_clone = Arc::clone(&received);
1434 session
1435 .bind_extensions(ExtensionBindings {
1436 on_error: Some(Arc::new(move |path, event, error| {
1437 received_clone
1438 .lock()
1439 .unwrap_or_else(std::sync::PoisonError::into_inner)
1440 .push((path.to_owned(), event.to_owned(), error.to_owned()));
1441 })),
1442 ..Default::default()
1443 })
1444 .await?;
1445 session.report_extension_error("/ext/tool.ts", "tool_call", "host crashed");
1446 assert_eq!(
1447 received
1448 .lock()
1449 .unwrap_or_else(std::sync::PoisonError::into_inner)
1450 .as_slice(),
1451 &[(
1452 "/ext/tool.ts".to_owned(),
1453 "tool_call".to_owned(),
1454 "host crashed".to_owned(),
1455 )]
1456 );
1457
1458 session
1459 .bind_extensions(ExtensionBindings {
1460 on_error: Some(Arc::new(|_, _, _| {
1461 std::panic::resume_unwind(Box::new("listener panic"));
1463 })),
1464 ..Default::default()
1465 })
1466 .await?;
1467 session.report_extension_error("<runtime>", "reload", "boom");
1468 assert!(!session.session_id().await.is_empty());
1469 Ok(())
1470 }
1471
1472 #[tokio::test]
1473 async fn bind_extensions_with_null_runner_is_noop() -> TestResult {
1474 let session = make_session()?;
1475 session
1476 .bind_extensions(ExtensionBindings {
1477 mode: Some(ExtensionMode::Print),
1478 ..Default::default()
1479 })
1480 .await?;
1481 assert!(!session.has_extension_handlers("session_start"));
1482 Ok(())
1483 }
1484
1485 #[tokio::test]
1486 async fn invoke_extension_shutdown_handler_calls_bound_closure() -> TestResult {
1487 let session = make_session()?;
1488 let called = Arc::new(AtomicBool::new(false));
1489 let called_clone = Arc::clone(&called);
1490 session
1491 .bind_extensions(ExtensionBindings {
1492 shutdown_handler: Some(Arc::new(move || {
1493 called_clone.store(true, Ordering::SeqCst);
1494 })),
1495 ..Default::default()
1496 })
1497 .await?;
1498 session.invoke_extension_shutdown_handler();
1499 assert!(called.load(Ordering::SeqCst));
1500 Ok(())
1501 }
1502
1503 enum FakeCmd {
1508 Emit(Frame),
1509 }
1510
1511 #[derive(Clone)]
1512 struct FakeHost {
1513 cmd_tx: mpsc::Sender<FakeCmd>,
1514 drop_methods: Arc<StdMutex<HashSet<String>>>,
1515 requests: Arc<StdMutex<Vec<Frame>>>,
1516 }
1517
1518 impl FakeHost {
1519 fn drop_method(&self, method: &str) {
1520 if let Ok(mut set) = self.drop_methods.lock() {
1521 set.insert(method.to_owned());
1522 }
1523 }
1524
1525 async fn emit(&self, frame: Frame) {
1526 let _ = self.cmd_tx.send(FakeCmd::Emit(frame)).await;
1527 }
1528
1529 async fn wait_for_request(&self, method: &str) -> TestResult {
1530 tokio::time::timeout(Duration::from_secs(1), async {
1531 loop {
1532 if self.requests.lock().is_ok_and(|requests| {
1533 requests.iter().any(|request| request.method == method)
1534 }) {
1535 return;
1536 }
1537 tokio::task::yield_now().await;
1538 }
1539 })
1540 .await
1541 .map_err(|_| io::Error::other(format!("fake host did not receive {method}")))?;
1542 Ok(())
1543 }
1544 }
1545
1546 fn last_request_id(host: &FakeHost, method: &str) -> TestResult<u64> {
1547 host.requests
1548 .lock()
1549 .map_err(|_| io::Error::other("request lock poisoned"))?
1550 .iter()
1551 .rev()
1552 .find(|request| request.method == method)
1553 .map(|request| request.id)
1554 .ok_or_else(|| io::Error::other(format!("no recorded {method} request")).into())
1555 }
1556
1557 fn recorded_request_count(host: &FakeHost) -> TestResult<usize> {
1558 Ok(host
1559 .requests
1560 .lock()
1561 .map_err(|_| io::Error::other("request lock poisoned"))?
1562 .len())
1563 }
1564
1565 fn cutover_snapshot(tool_name: &str, provider_name: &str) -> Value {
1566 json!({
1567 "tools": [
1568 {
1569 "name": tool_name,
1570 "label": tool_name,
1571 "description": "cutover tool",
1572 "parameters": {"type": "object"}
1573 }
1574 ],
1575 "commands": [],
1576 "shortcuts": [],
1577 "flags": [{"name": "extFlag", "type": "string", "default": "x"}],
1578 "renderers": [],
1579 "providers": [{"name": provider_name}],
1580 "handlers": [
1581 "session_start",
1582 "session_shutdown",
1583 "resources_discover",
1584 "tool_call"
1585 ],
1586 })
1587 }
1588
1589 fn dispatch(
1590 req: &Frame,
1591 snapshot: &Value,
1592 responses: &StdMutex<HashMap<String, Value>>,
1593 drop_methods: &StdMutex<HashSet<String>>,
1594 ) -> Option<Frame> {
1595 if drop_methods
1596 .lock()
1597 .is_ok_and(|set| set.contains(&req.method))
1598 {
1599 return None;
1600 }
1601 let payload = if req.method == "hello" {
1602 serde_json::to_value(HelloAck::local()).unwrap_or(Value::Null)
1603 } else if req.method == "extensions.load" {
1604 snapshot.clone()
1605 } else if let Some(payload) = responses
1606 .lock()
1607 .ok()
1608 .and_then(|map| map.get(&req.method).cloned())
1609 {
1610 payload
1611 } else if req.method == pi_ext::protocol::FLAGS_SET_METHOD {
1612 json!({"ok": true})
1613 } else {
1614 Value::Object(Map::new())
1615 };
1616 Some(Frame {
1617 id: req.id,
1618 kind: FrameKind::Res,
1619 method: req.method.clone(),
1620 payload,
1621 })
1622 }
1623
1624 async fn fake_host_task(
1625 read: tokio::io::DuplexStream,
1626 mut write: tokio::io::DuplexStream,
1627 snapshot: Value,
1628 responses: Arc<StdMutex<HashMap<String, Value>>>,
1629 drop_methods: Arc<StdMutex<HashSet<String>>>,
1630 requests: Arc<StdMutex<Vec<Frame>>>,
1631 mut cmd_rx: mpsc::Receiver<FakeCmd>,
1632 ) {
1633 let mut reader = BufReader::new(read);
1634 let mut buf = String::new();
1635 loop {
1636 tokio::select! {
1637 biased;
1638 cmd = cmd_rx.recv() => {
1639 match cmd {
1640 Some(FakeCmd::Emit(frame)) => {
1641 let bytes = encode_frame(&frame).unwrap_or_default();
1642 if !bytes.is_empty() {
1643 let _ = write.write_all(&bytes).await;
1644 let _ = write.flush().await;
1645 }
1646 }
1647 None => return,
1648 }
1649 }
1650 n = reader.read_line(&mut buf) => {
1651 match n {
1652 Ok(0) | Err(_) => return,
1653 Ok(_) => {
1654 if let Ok(req) = decode_frame_str(&buf) {
1655 if let Ok(mut recorded) = requests.lock() {
1656 recorded.push(req.clone());
1657 }
1658 if let Some(resp) =
1659 dispatch(&req, &snapshot, &responses, &drop_methods)
1660 {
1661 let bytes = encode_frame(&resp).unwrap_or_default();
1662 let _ = write.write_all(&bytes).await;
1663 let _ = write.flush().await;
1664 }
1665 }
1666 buf.clear();
1667 }
1668 }
1669 }
1670 }
1671 }
1672 }
1673
1674 async fn make_host_runner(
1675 snapshot: Value,
1676 hook_timeout: Duration,
1677 ) -> TestResult<(Arc<HostExtensionRunner>, FakeHost)> {
1678 let (client_to_host, host_read) = tokio::io::duplex(64 * 1024);
1679 let (host_write, client_read) = tokio::io::duplex(64 * 1024);
1680 let (err_write, _err_read) = tokio::io::duplex(4096);
1681 let client = Arc::new(HostClient::connect_boxed(
1682 Box::new(client_to_host),
1683 Box::new(client_read),
1684 Box::new(err_write),
1685 None,
1686 ));
1687 let responses = Arc::new(StdMutex::new(HashMap::new()));
1688 let drop_methods = Arc::new(StdMutex::new(HashSet::new()));
1689 let requests = Arc::new(StdMutex::new(Vec::new()));
1690 let (cmd_tx, cmd_rx) = mpsc::channel(64);
1691 tokio::spawn(fake_host_task(
1692 host_read,
1693 host_write,
1694 snapshot,
1695 Arc::clone(&responses),
1696 Arc::clone(&drop_methods),
1697 Arc::clone(&requests),
1698 cmd_rx,
1699 ));
1700 let runner = HostExtensionRunner::connect_with_cwd_and_trust(
1701 client,
1702 vec![],
1703 "/workspace",
1704 false,
1705 hook_timeout,
1706 )
1707 .await?;
1708 Ok((
1709 runner,
1710 FakeHost {
1711 cmd_tx,
1712 drop_methods,
1713 requests,
1714 },
1715 ))
1716 }
1717
1718 type ReplacementHostSlot = Arc<StdMutex<Option<(Arc<HostExtensionRunner>, FakeHost)>>>;
1719
1720 struct ReloadCutoverFixture {
1721 runtime: Arc<ModelRuntime>,
1722 old_host: Arc<HostExtensionRunner>,
1723 old_fake: FakeHost,
1724 session: Arc<AgentSession>,
1725 replacement: ReplacementHostSlot,
1726 }
1727
1728 async fn reload_cutover_fixture() -> TestResult<ReloadCutoverFixture> {
1729 let runtime = Arc::new(ModelRuntime::create_in_memory().await?);
1730 let (old_host, old_fake) = make_host_runner(
1731 cutover_snapshot("oldTool", "oldProv"),
1732 Duration::from_secs(5),
1733 )
1734 .await?;
1735 let registration = old_host.register_providers_on(runtime.as_ref());
1736 assert!(registration.iter().all(|(_, result)| result.is_ok()));
1737 assert!(
1738 runtime.get_registered_provider_config("oldProv").is_some(),
1739 "old provider must be registered before reload"
1740 );
1741
1742 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1743 config.extension_runner = Some(Arc::clone(&old_host) as Arc<dyn ExtensionRunner>);
1744 config.host_extension_runner = Some(Arc::clone(&old_host));
1745 config.model_runtime = Some(Arc::clone(&runtime));
1746 config.initial_active_tool_names = Some(vec!["oldTool".to_owned()]);
1747 let session = AgentSession::new(config)?;
1748 session.refresh_tool_registry(&super::super::tools::RefreshToolRegistryOptions {
1749 active_tool_names: Some(vec!["oldTool".to_owned()]),
1750 include_all_extension_tools: true,
1751 });
1752 assert!(
1753 session.get_tool("oldTool").is_some(),
1754 "session registry must start with the old host tool"
1755 );
1756
1757 let replacement = Arc::new(StdMutex::new(None));
1758 let replacement_for_factory = Arc::clone(&replacement);
1759 session.set_reload_restart_factory(Some(Arc::new(move |_paths, _cwd, _trusted| {
1760 let replacement = Arc::clone(&replacement_for_factory);
1761 Box::pin(async move {
1762 let (runner, host) = make_host_runner(
1763 cutover_snapshot("newTool", "newProv"),
1764 Duration::from_secs(5),
1765 )
1766 .await
1767 .map_err(|error| HostStartError::Load(error.to_string()))?;
1768 *replacement
1769 .lock()
1770 .map_err(|_| HostStartError::Load("replacement lock poisoned".to_owned()))? =
1771 Some((Arc::clone(&runner), host));
1772 Ok(runner)
1773 })
1774 })));
1775
1776 Ok(ReloadCutoverFixture {
1777 runtime,
1778 old_host,
1779 old_fake,
1780 session,
1781 replacement,
1782 })
1783 }
1784
1785 async fn wait_for_session_cutover(fixture: &ReloadCutoverFixture) -> TestResult {
1786 tokio::time::timeout(Duration::from_secs(2), async {
1787 loop {
1788 let host = fixture.session.host_extension_runner();
1789 let cut_over = host
1790 .as_ref()
1791 .is_some_and(|host| !Arc::ptr_eq(host, &fixture.old_host))
1792 && fixture.session.get_tool("newTool").is_some()
1793 && fixture.session.get_tool("oldTool").is_none()
1794 && fixture
1795 .runtime
1796 .get_registered_provider_config("newProv")
1797 .is_some()
1798 && fixture
1799 .runtime
1800 .get_registered_provider_config("oldProv")
1801 .is_none()
1802 && fixture.old_host.is_running();
1803 if cut_over {
1804 break;
1805 }
1806 tokio::task::yield_now().await;
1807 }
1808 })
1809 .await
1810 .map_err(|_| {
1811 io::Error::other("session surfaces never cut over while old host stayed live")
1812 })?;
1813 Ok(())
1814 }
1815
1816 async fn verify_replacement_lifecycle(fixture: &ReloadCutoverFixture) -> TestResult {
1817 assert!(
1818 !fixture.old_host.is_running(),
1819 "old host must be reaped only after cutover completes"
1820 );
1821 assert!(
1822 fixture
1823 .session
1824 .host_extension_runner()
1825 .is_some_and(|host| host.is_running()),
1826 "replacement host remains live after reload"
1827 );
1828 assert!(
1829 fixture.session.get_tool("newTool").is_some(),
1830 "tool registry must retain the replacement tool"
1831 );
1832 assert!(
1833 fixture
1834 .runtime
1835 .get_registered_provider_config("newProv")
1836 .is_some(),
1837 "runtime must retain the replacement provider"
1838 );
1839
1840 let old_methods: Vec<String> = fixture
1841 .old_fake
1842 .requests
1843 .lock()
1844 .map_err(|_| io::Error::other("old request lock poisoned"))?
1845 .iter()
1846 .map(|frame| frame.method.clone())
1847 .collect();
1848 assert!(
1849 old_methods
1850 .iter()
1851 .any(|method| method == "session_shutdown"),
1852 "old host must observe session_shutdown before cutover: {old_methods:?}"
1853 );
1854 let (replacement_runner, replacement_fake) = fixture
1855 .replacement
1856 .lock()
1857 .map_err(|_| io::Error::other("replacement lock poisoned"))?
1858 .clone()
1859 .ok_or("replacement runner missing")?;
1860 assert!(
1861 fixture
1862 .session
1863 .host_extension_runner()
1864 .is_some_and(|host| Arc::ptr_eq(&host, &replacement_runner)),
1865 "session host handle must be the factory-produced replacement"
1866 );
1867 let new_methods: Vec<String> = replacement_fake
1868 .requests
1869 .lock()
1870 .map_err(|_| io::Error::other("replacement request lock poisoned"))?
1871 .iter()
1872 .map(|frame| frame.method.clone())
1873 .collect();
1874 let start_idx = new_methods
1875 .iter()
1876 .position(|method| method == "session_start")
1877 .ok_or("replacement must receive session_start")?;
1878 let discover_idx = new_methods
1879 .iter()
1880 .position(|method| method == "resources_discover")
1881 .ok_or("replacement must receive resources_discover")?;
1882 assert!(
1883 start_idx < discover_idx,
1884 "lifecycle must be start then discover on the replacement: {new_methods:?}"
1885 );
1886 replacement_runner.shutdown_once().await;
1887 Ok(())
1888 }
1889
1890 #[tokio::test]
1891 async fn reload_cuts_over_session_surfaces_before_reaping_old_host() -> TestResult {
1892 let fixture = reload_cutover_fixture().await?;
1893
1894 fixture.old_fake.drop_method("tool_call");
1895 let blocked = {
1896 let old_host = Arc::clone(&fixture.old_host);
1897 tokio::spawn(async move {
1898 ExtensionRunner::emit_tool_call(
1899 old_host.as_ref(),
1900 "read",
1901 "tc-session-cutover",
1902 Map::new(),
1903 )
1904 .await
1905 })
1906 };
1907 fixture.old_fake.wait_for_request("tool_call").await?;
1908
1909 let reloading = {
1910 let session = Arc::clone(&fixture.session);
1911 tokio::spawn(async move { session.reload().await })
1912 };
1913 wait_for_session_cutover(&fixture).await?;
1914
1915 assert!(
1916 !reloading.is_finished(),
1917 "reload must stay pending until blocked old traffic drains"
1918 );
1919 assert!(
1920 fixture.old_host.is_running(),
1921 "old transport must remain live until after session cutover"
1922 );
1923 let new_host = fixture
1924 .session
1925 .host_extension_runner()
1926 .ok_or("missing replacement host handle")?;
1927 assert!(
1928 !Arc::ptr_eq(&new_host, &fixture.old_host),
1929 "host handle must point at the replacement"
1930 );
1931 assert!(
1932 Arc::ptr_eq(
1933 &fixture.session.extension_runner(),
1934 &(Arc::clone(&new_host) as Arc<dyn ExtensionRunner>)
1935 ) || fixture
1936 .session
1937 .extension_runner()
1938 .has_handlers("session_start"),
1939 "trait runner must be the replacement after cutover"
1940 );
1941 assert!(
1942 fixture
1943 .session
1944 .extension_runner()
1945 .has_handlers("session_start"),
1946 "post-cutover trait runner must expose replacement handlers"
1947 );
1948
1949 let id = last_request_id(&fixture.old_fake, "tool_call")?;
1950 fixture
1951 .old_fake
1952 .emit(Frame {
1953 id,
1954 kind: FrameKind::Res,
1955 method: "tool_call".to_owned(),
1956 payload: json!({"block": true, "reason": "session-cutover"}),
1957 })
1958 .await;
1959 let verdict = tokio::time::timeout(Duration::from_secs(2), blocked)
1960 .await???
1961 .ok_or("blocked old hook must resolve against the still-live transport")?;
1962 assert!(
1963 verdict.block,
1964 "old host verdict must survive session cutover"
1965 );
1966
1967 tokio::time::timeout(Duration::from_secs(2), reloading)
1968 .await??
1969 .map_err(|error| io::Error::other(error.to_string()))?;
1970 verify_replacement_lifecycle(&fixture).await
1971 }
1972
1973 #[tokio::test]
1974 async fn reload_keeps_old_surfaces_when_replacement_start_fails() -> TestResult {
1975 let runtime = Arc::new(ModelRuntime::create_in_memory().await?);
1976 let (old_host, _old_fake) = make_host_runner(
1977 cutover_snapshot("oldTool", "oldProv"),
1978 Duration::from_secs(1),
1979 )
1980 .await?;
1981 let registration = old_host.register_providers_on(runtime.as_ref());
1982 assert!(registration.iter().all(|(_, result)| result.is_ok()));
1983
1984 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1985 config.extension_runner = Some(Arc::clone(&old_host) as Arc<dyn ExtensionRunner>);
1986 config.host_extension_runner = Some(Arc::clone(&old_host));
1987 config.model_runtime = Some(Arc::clone(&runtime));
1988 config.initial_active_tool_names = Some(vec!["oldTool".to_owned()]);
1989 let session = AgentSession::new(config)?;
1990 session.refresh_tool_registry(&super::super::tools::RefreshToolRegistryOptions {
1991 active_tool_names: Some(vec!["oldTool".to_owned()]),
1992 include_all_extension_tools: true,
1993 });
1994
1995 session.set_reload_restart_factory(Some(Arc::new(|_paths, _cwd, _trusted| {
1996 Box::pin(async {
1997 Err(HostStartError::Spawn(
1998 "injected replacement start failure".to_owned(),
1999 ))
2000 })
2001 })));
2002
2003 let error = match session.reload().await {
2004 Ok(()) => return Err("failed replacement start unexpectedly succeeded".into()),
2005 Err(error) => error,
2006 };
2007 assert!(
2008 matches!(error, ExtensionBindError::HostRestart(_)),
2009 "reload must surface host restart failure: {error}"
2010 );
2011 assert!(
2012 session
2013 .host_extension_runner()
2014 .is_some_and(|host| Arc::ptr_eq(&host, &old_host)),
2015 "failed reload must keep the old host handle"
2016 );
2017 assert!(
2018 old_host.is_running(),
2019 "failed reload must not reap the old host"
2020 );
2021 assert!(
2022 session.get_tool("oldTool").is_some(),
2023 "failed reload must keep the old tool registry"
2024 );
2025 assert!(
2026 runtime.get_registered_provider_config("oldProv").is_some(),
2027 "failed reload must keep the old provider registration"
2028 );
2029
2030 old_host.shutdown_once().await;
2031 Ok(())
2032 }
2033
2034 #[tokio::test]
2035 async fn reload_reaping_before_cutover_would_fail_blocked_old_traffic() -> TestResult {
2036 let runtime = Arc::new(ModelRuntime::create_in_memory().await?);
2039 let (old_host, old_fake) = make_host_runner(
2040 cutover_snapshot("oldTool", "oldProv"),
2041 Duration::from_secs(5),
2042 )
2043 .await?;
2044 let _ = old_host.register_providers_on(runtime.as_ref());
2045
2046 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
2047 config.extension_runner = Some(Arc::clone(&old_host) as Arc<dyn ExtensionRunner>);
2048 config.host_extension_runner = Some(Arc::clone(&old_host));
2049 config.model_runtime = Some(Arc::clone(&runtime));
2050 let session = AgentSession::new(config)?;
2051
2052 let replacement_keep = Arc::new(StdMutex::new(None::<FakeHost>));
2053 let replacement_keep_for_factory = Arc::clone(&replacement_keep);
2054 let cutover_seen = Arc::new(AtomicBool::new(false));
2055 let cutover_seen_for_factory = Arc::clone(&cutover_seen);
2056 let old_host_for_factory = Arc::clone(&old_host);
2057 session.set_reload_restart_factory(Some(Arc::new(move |_paths, _cwd, _trusted| {
2058 let cutover_seen = Arc::clone(&cutover_seen_for_factory);
2059 let old_host = Arc::clone(&old_host_for_factory);
2060 let replacement_keep = Arc::clone(&replacement_keep_for_factory);
2061 Box::pin(async move {
2062 assert!(
2065 old_host.is_running(),
2066 "replacement factory must run while old host is still live"
2067 );
2068 let (runner, host) = make_host_runner(
2069 cutover_snapshot("newTool", "newProv"),
2070 Duration::from_secs(5),
2071 )
2072 .await
2073 .map_err(|error| HostStartError::Load(error.to_string()))?;
2074 *replacement_keep.lock().map_err(|_| {
2075 HostStartError::Load("replacement keep lock poisoned".to_owned())
2076 })? = Some(host);
2077 cutover_seen.store(true, Ordering::SeqCst);
2078 Ok(runner)
2079 })
2080 })));
2081
2082 old_fake.drop_method("tool_call");
2083 let frames_before = recorded_request_count(&old_fake)?;
2084 let blocked = {
2085 let old_host = Arc::clone(&old_host);
2086 tokio::spawn(async move {
2087 ExtensionRunner::emit_tool_call(
2088 old_host.as_ref(),
2089 "read",
2090 "tc-mutation",
2091 Map::new(),
2092 )
2093 .await
2094 })
2095 };
2096 old_fake.wait_for_request("tool_call").await?;
2097 assert!(
2098 recorded_request_count(&old_fake)? > frames_before,
2099 "blocked hook must reach the old transport"
2100 );
2101
2102 let reloading = {
2103 let session = Arc::clone(&session);
2104 tokio::spawn(async move { session.reload().await })
2105 };
2106
2107 tokio::time::timeout(Duration::from_secs(2), async {
2110 loop {
2111 if cutover_seen.load(Ordering::SeqCst) && old_host.is_running() {
2112 break;
2113 }
2114 tokio::task::yield_now().await;
2115 }
2116 })
2117 .await
2118 .map_err(|_| io::Error::other("replacement never prepared while old host stayed live"))?;
2119
2120 assert!(!blocked.is_finished(), "old hook must still be in flight");
2123 assert!(!reloading.is_finished(), "reload must wait for old drain");
2124
2125 let id = last_request_id(&old_fake, "tool_call")?;
2126 old_fake
2127 .emit(Frame {
2128 id,
2129 kind: FrameKind::Res,
2130 method: "tool_call".to_owned(),
2131 payload: json!({"block": false}),
2132 })
2133 .await;
2134 let _ = tokio::time::timeout(Duration::from_secs(2), blocked).await???;
2135 tokio::time::timeout(Duration::from_secs(2), reloading)
2136 .await??
2137 .map_err(|error| io::Error::other(error.to_string()))?;
2138 assert!(!old_host.is_running(), "old host reaped after drain");
2139 let _ = replacement_keep
2141 .lock()
2142 .map_err(|_| io::Error::other("replacement keep lock poisoned"))?
2143 .take();
2144 Ok(())
2145 }
2146}