1use std::collections::{BTreeMap, BTreeSet};
8use std::fmt::Write as _;
9use std::path::{Path, PathBuf};
10use std::sync::Arc;
11
12use pi_ai::providers::KnownProvider;
13use pi_ai::types::{Model, ModelThinkingLevel};
14use pi_ext::protocol::FlagValueWire;
15use thiserror::Error;
16
17use super::config::{get_agent_dir, get_docs_path, resolve_path};
18use super::extension_host::HostExtensionRunner;
19use super::model_runtime::{
20 CreateModelRuntimeOptions, ModelRuntime, ModelRuntimeError, ProviderConfigInput,
21};
22use super::resources::{DefaultResourceLoader, DefaultResourceLoaderOptions, ResourceLoader};
23use super::settings::{SettingsManager, SettingsManagerCreateOptions};
24use super::trust::{ProjectTrustStore, ResolveProjectTrustedOptions, resolve_project_trusted};
25
26pub const DEFAULT_THINKING_LEVEL: ModelThinkingLevel = ModelThinkingLevel::Medium;
28
29#[must_use]
31pub fn default_model_per_provider() -> &'static [(&'static str, &'static str)] {
32 &[
33 ("amazon-bedrock", "us.anthropic.claude-opus-4-6-v1"),
34 ("ant-ling", "Ring-2.6-1T"),
35 ("anthropic", "claude-opus-4-8"),
36 ("openai", "gpt-5.5"),
37 ("azure-openai-responses", "gpt-5.4"),
38 ("openai-codex", "gpt-5.5"),
39 ("radius", "auto"),
40 ("nvidia", "nvidia/nemotron-3-super-120b-a12b"),
41 ("deepseek", "deepseek-v4-pro"),
42 ("google", "gemini-3.1-pro-preview"),
43 ("google-vertex", "gemini-3.1-pro-preview"),
44 ("github-copilot", "gpt-5.4"),
45 ("openrouter", "moonshotai/kimi-k2.6"),
46 ("vercel-ai-gateway", "zai/glm-5.1"),
47 ("xai", "grok-4.5"),
48 ("groq", "openai/gpt-oss-120b"),
49 ("cerebras", "zai-glm-4.7"),
50 ("zai", "glm-5.1"),
51 ("zai-coding-cn", "glm-5.1"),
52 ("mistral", "devstral-medium-latest"),
53 ("minimax", "MiniMax-M2.7"),
54 ("minimax-cn", "MiniMax-M2.7"),
55 ("moonshotai", "kimi-k2.6"),
56 ("moonshotai-cn", "kimi-k2.6"),
57 ("huggingface", "moonshotai/Kimi-K2.6"),
58 ("fireworks", "accounts/fireworks/models/kimi-k2p6"),
59 ("together", "moonshotai/Kimi-K2.6"),
60 ("opencode", "kimi-k2.6"),
61 ("opencode-go", "kimi-k2.6"),
62 ("kimi-coding", "kimi-for-coding"),
63 ("cloudflare-workers-ai", "@cf/moonshotai/kimi-k2.6"),
64 (
65 "cloudflare-ai-gateway",
66 "workers-ai/@cf/moonshotai/kimi-k2.6",
67 ),
68 ("xiaomi", "mimo-v2.5-pro"),
69 ("xiaomi-token-plan-cn", "mimo-v2.5-pro"),
70 ("xiaomi-token-plan-ams", "mimo-v2.5-pro"),
71 ("xiaomi-token-plan-sgp", "mimo-v2.5-pro"),
72 ]
73}
74
75#[derive(Clone, Copy, Debug, Eq, PartialEq)]
77pub enum AgentSessionRuntimeDiagnosticKind {
78 Info,
80 Warning,
82 Error,
84}
85
86#[derive(Clone, Debug, Eq, PartialEq)]
92pub struct AgentSessionRuntimeDiagnostic {
93 pub kind: AgentSessionRuntimeDiagnosticKind,
95 pub message: String,
97}
98
99impl AgentSessionRuntimeDiagnostic {
100 #[must_use]
102 pub fn info(message: impl Into<String>) -> Self {
103 Self {
104 kind: AgentSessionRuntimeDiagnosticKind::Info,
105 message: message.into(),
106 }
107 }
108
109 #[must_use]
111 pub fn warning(message: impl Into<String>) -> Self {
112 Self {
113 kind: AgentSessionRuntimeDiagnosticKind::Warning,
114 message: message.into(),
115 }
116 }
117
118 #[must_use]
120 pub fn error(message: impl Into<String>) -> Self {
121 Self {
122 kind: AgentSessionRuntimeDiagnosticKind::Error,
123 message: message.into(),
124 }
125 }
126}
127
128#[derive(Default)]
130pub struct CreateAgentSessionServicesOptions {
131 pub cwd: PathBuf,
133 pub agent_dir: Option<PathBuf>,
135 pub settings_manager: Option<SettingsManager>,
137 pub model_runtime: Option<ModelRuntime>,
139 pub extension_flag_values: Option<BTreeMap<String, ExtensionFlagValue>>,
141 pub resource_loader_options: Option<ResourceLoaderServiceOptions>,
143 pub pending_provider_registrations: Vec<PendingProviderRegistration>,
149 pub registered_extension_flags: BTreeMap<String, ExtensionFlagType>,
153}
154
155pub type ResourceDiscoveryDisabled = bool;
160
161#[derive(Clone, Debug, Default)]
163pub struct ResourceLoaderServiceOptions {
164 pub additional_extension_paths: Vec<String>,
166 pub additional_skill_paths: Vec<String>,
168 pub additional_prompt_template_paths: Vec<String>,
170 pub additional_theme_paths: Vec<String>,
172 pub no_extensions: ResourceDiscoveryDisabled,
174 pub no_skills: ResourceDiscoveryDisabled,
176 pub no_prompt_templates: ResourceDiscoveryDisabled,
178 pub no_themes: ResourceDiscoveryDisabled,
180 pub no_context_files: ResourceDiscoveryDisabled,
182 pub system_prompt: Option<String>,
184 pub append_system_prompt: Option<Vec<String>>,
186}
187
188#[derive(Clone, Copy, Debug, Default)]
189struct ResourceDiscoveryPolicy(u8);
190
191impl ResourceDiscoveryPolicy {
192 const EXTENSIONS: u8 = 1 << 0;
193 const SKILLS: u8 = 1 << 1;
194 const PROMPT_TEMPLATES: u8 = 1 << 2;
195 const THEMES: u8 = 1 << 3;
196 const CONTEXT_FILES: u8 = 1 << 4;
197
198 fn from_options(options: &ResourceLoaderServiceOptions) -> Self {
199 let mut bits = 0;
200 for (disabled, flag) in [
201 (options.no_extensions, Self::EXTENSIONS),
202 (options.no_skills, Self::SKILLS),
203 (options.no_prompt_templates, Self::PROMPT_TEMPLATES),
204 (options.no_themes, Self::THEMES),
205 (options.no_context_files, Self::CONTEXT_FILES),
206 ] {
207 if disabled {
208 bits |= flag;
209 }
210 }
211 Self(bits)
212 }
213
214 const fn disables(self, flag: u8) -> bool {
215 self.0 & flag != 0
216 }
217}
218
219#[derive(Clone, Debug)]
221pub struct PendingProviderRegistration {
222 pub name: String,
224 pub config: ProviderConfigInput,
226 pub extension_path: String,
228}
229
230#[derive(Clone, Copy, Debug, Eq, PartialEq)]
232pub enum ExtensionFlagType {
233 Boolean,
235 String,
237}
238
239#[derive(Clone, Debug, Eq, PartialEq)]
241pub enum ExtensionFlagValue {
242 Bool(bool),
244 Str(String),
246}
247
248pub struct AgentSessionServices {
253 pub cwd: PathBuf,
255 pub agent_dir: PathBuf,
257 pub model_runtime: ModelRuntime,
259 pub resource_loader: DefaultResourceLoader,
261 pub diagnostics: Vec<AgentSessionRuntimeDiagnostic>,
263 pub extension_flag_values: BTreeMap<String, ExtensionFlagValue>,
265 pub extension_runner: Option<Arc<HostExtensionRunner>>,
270}
271
272impl AgentSessionServices {
273 #[must_use]
275 pub fn settings_manager(&self) -> &SettingsManager {
276 self.resource_loader.settings_manager()
277 }
278
279 pub fn settings_manager_mut(&mut self) -> &mut SettingsManager {
281 self.resource_loader.settings_manager_mut()
282 }
283}
284
285pub struct CreateAgentSessionFromServicesOptions {
287 pub services: AgentSessionServices,
289 pub model: Option<Model>,
291 pub thinking_level: Option<ModelThinkingLevel>,
293 pub scoped_models: Vec<ScopedModel>,
295 pub tools: Option<Vec<String>>,
297 pub exclude_tools: Option<Vec<String>>,
299 pub no_tools: Option<NoToolsMode>,
301 pub session_start_event: Option<crate::core::agent_session::SessionStartEvent>,
304 pub saved_session_model: Option<(String, String)>,
306 pub has_existing_session: bool,
308}
309
310#[derive(Clone, Copy, Debug, Eq, PartialEq)]
312pub enum NoToolsMode {
313 All,
315 Builtin,
317}
318
319#[derive(Clone, Debug, PartialEq)]
321pub struct ScopedModel {
322 pub model: Model,
324 pub thinking_level: Option<ModelThinkingLevel>,
326}
327
328#[derive(Clone, Debug, PartialEq)]
330pub struct InitialModelResult {
331 pub model: Option<Model>,
333 pub thinking_level: ModelThinkingLevel,
335 pub fallback_message: Option<String>,
337}
338
339pub struct CreateAgentSessionResult {
347 pub model: Option<Model>,
349 pub thinking_level: ModelThinkingLevel,
351 pub initial_active_tool_names: Vec<String>,
353 pub allowed_tool_names: Option<Vec<String>>,
355 pub excluded_tool_names: Option<Vec<String>>,
357 pub scoped_models: Vec<ScopedModel>,
359 pub model_fallback_message: Option<String>,
361 pub diagnostics: Vec<AgentSessionRuntimeDiagnostic>,
363 pub cwd: PathBuf,
365 pub agent_dir: PathBuf,
367 pub model_runtime: ModelRuntime,
369 pub resource_loader: DefaultResourceLoader,
371 pub session_start_event: Option<crate::core::agent_session::SessionStartEvent>,
373 pub extension_runner: Option<Arc<HostExtensionRunner>>,
375}
376
377#[derive(Clone, Debug, Error)]
379pub enum AgentSessionServicesError {
380 #[error(transparent)]
382 ModelRuntime(#[from] ModelRuntimeError),
383 #[error("{0}")]
385 Trust(String),
386 #[error("{0}")]
388 ResourceLoader(String),
389}
390
391#[must_use]
393pub fn get_provider_login_help() -> String {
394 let docs = get_docs_path();
395 format!(
396 "Use /login to log into a provider via OAuth or API key. See:\n {}\n {}",
397 docs.join("providers.md").display(),
398 docs.join("models.md").display()
399 )
400}
401
402#[must_use]
404pub fn format_no_models_available_message() -> String {
405 format!("No models available. {}", get_provider_login_help())
406}
407
408#[must_use]
410pub fn format_no_model_selected_message() -> String {
411 format!(
412 "No model selected.\n\n{}\n\nThen use /model to select a model.",
413 get_provider_login_help()
414 )
415}
416
417#[must_use]
419pub fn format_no_api_key_found_message(provider: &str) -> String {
420 let provider_display = if provider == "unknown" {
421 "the selected model"
422 } else {
423 provider
424 };
425 format!(
426 "No API key found for {provider_display}.\n\n{}",
427 get_provider_login_help()
428 )
429}
430
431#[must_use]
433pub fn format_oauth_auth_failed_message(provider: &str) -> String {
434 format!(
435 "Authentication failed for \"{provider}\". Credentials may have expired or network is unavailable. Run '/login {provider}' to re-authenticate."
436 )
437}
438
439pub async fn create_agent_session_services(
445 options: CreateAgentSessionServicesOptions,
446) -> Result<AgentSessionServices, AgentSessionServicesError> {
447 create_agent_session_services_with_trust(options, None).await
448}
449
450pub async fn create_agent_session_services_with_trust(
458 options: CreateAgentSessionServicesOptions,
459 project_trust_override: Option<bool>,
460) -> Result<AgentSessionServices, AgentSessionServicesError> {
461 let CreateAgentSessionServicesOptions {
462 cwd,
463 agent_dir,
464 settings_manager,
465 model_runtime,
466 extension_flag_values,
467 resource_loader_options,
468 pending_provider_registrations,
469 registered_extension_flags,
470 } = options;
471 let settings_manager_was_supplied = settings_manager.is_some();
472 let mut foundation =
473 create_service_foundation(cwd, agent_dir, settings_manager, model_runtime).await?;
474 let project_trusted = if settings_manager_was_supplied && project_trust_override.is_none() {
475 foundation.settings_manager.is_project_trusted()
476 } else {
477 resolve_project_trusted(ResolveProjectTrustedOptions {
478 cwd: foundation.cwd.clone(),
479 trust_store: &ProjectTrustStore::new(&foundation.agent_dir),
480 trust_override: project_trust_override,
481 default_project_trust: foundation.settings_manager.get_default_project_trust(),
482 extension_hook: None,
483 ui: None,
484 on_extension_error: None,
485 })
486 .map_err(|error| AgentSessionServicesError::Trust(error.to_string()))?
487 };
488 foundation
489 .settings_manager
490 .set_project_trusted(project_trusted);
491 let (resource_loader, discovery) = create_service_resource_loader(
492 &foundation.cwd,
493 &foundation.agent_dir,
494 foundation.settings_manager,
495 resource_loader_options.unwrap_or_default(),
496 )
497 .await?;
498
499 let mut diagnostics = extension_discovery_diagnostics(&resource_loader);
500 let (mut extension_runner, host_registered_flags) = start_extension_phase(
501 &resource_loader,
502 discovery,
503 &foundation.cwd,
504 &foundation.model_runtime,
505 project_trusted,
506 &mut diagnostics,
507 )
508 .await;
509
510 let mut registered_flags = registered_extension_flags;
511 registered_flags.extend(host_registered_flags);
512 let (flag_diagnostics, applied_flags) =
513 apply_extension_flag_values(extension_flag_values.unwrap_or_default(), ®istered_flags);
514 diagnostics.extend(flag_diagnostics);
515 if let Some(runner) = extension_runner.as_deref()
516 && let Err(error) = apply_flags_to_runner(runner, &applied_flags).await
517 {
518 diagnostics.push(AgentSessionRuntimeDiagnostic::error(format!(
519 "Extension flags failed to apply: {error}"
520 )));
521 runner.unregister_providers_from(&foundation.model_runtime);
522 runner.shutdown_once().await;
523 extension_runner = None;
524 }
525
526 for registration in pending_provider_registrations {
527 if let Err(error) = foundation
528 .model_runtime
529 .register_provider(®istration.name, registration.config)
530 {
531 diagnostics.push(AgentSessionRuntimeDiagnostic::error(format!(
532 "Extension \"{}\" error: {error}",
533 registration.extension_path
534 )));
535 }
536 }
537
538 let _ = foundation
539 .model_runtime
540 .refresh(super::model_runtime::ModelsRefreshOptions {
541 allow_network: Some(false),
542 })
543 .await;
544
545 Ok(AgentSessionServices {
546 cwd: foundation.cwd,
547 agent_dir: foundation.agent_dir,
548 model_runtime: foundation.model_runtime,
549 resource_loader,
550 diagnostics,
551 extension_flag_values: applied_flags,
552 extension_runner,
553 })
554}
555
556struct ServiceFoundation {
557 cwd: PathBuf,
558 agent_dir: PathBuf,
559 model_runtime: ModelRuntime,
560 settings_manager: SettingsManager,
561}
562
563async fn create_service_foundation(
564 cwd: PathBuf,
565 agent_dir: Option<PathBuf>,
566 settings_manager: Option<SettingsManager>,
567 model_runtime: Option<ModelRuntime>,
568) -> Result<ServiceFoundation, AgentSessionServicesError> {
569 let cwd = resolve_path(cwd.to_string_lossy().as_ref());
570 let agent_dir = agent_dir.map_or_else(get_agent_dir, |path| {
571 resolve_path(path.to_string_lossy().as_ref())
572 });
573 let model_runtime = match model_runtime {
574 Some(runtime) => runtime,
575 None => {
576 ModelRuntime::create(CreateModelRuntimeOptions {
577 auth_path: Some(agent_dir.join("auth.json")),
578 models_path: Some(agent_dir.join("models.json")),
579 models_store_path: Some(agent_dir.join("models-store.json")),
580 allow_model_network: Some(false),
581 ..CreateModelRuntimeOptions::default()
582 })
583 .await?
584 }
585 };
586 let settings_manager = settings_manager.unwrap_or_else(|| {
587 SettingsManager::create(
588 &cwd,
589 Some(&agent_dir),
590 SettingsManagerCreateOptions::default().project_trusted(false),
591 )
592 });
593 Ok(ServiceFoundation {
594 cwd,
595 agent_dir,
596 model_runtime,
597 settings_manager,
598 })
599}
600
601async fn create_service_resource_loader(
602 cwd: &Path,
603 agent_dir: &Path,
604 settings_manager: SettingsManager,
605 options: ResourceLoaderServiceOptions,
606) -> Result<(DefaultResourceLoader, ResourceDiscoveryPolicy), AgentSessionServicesError> {
607 let discovery = ResourceDiscoveryPolicy::from_options(&options);
608 let mut loader = DefaultResourceLoader::new(DefaultResourceLoaderOptions {
609 cwd: cwd.to_path_buf(),
610 agent_dir: agent_dir.to_path_buf(),
611 settings_manager: Some(settings_manager),
612 additional_extension_paths: options.additional_extension_paths,
613 additional_skill_paths: options.additional_skill_paths,
614 additional_prompt_template_paths: options.additional_prompt_template_paths,
615 additional_theme_paths: options.additional_theme_paths,
616 no_extensions: discovery.disables(ResourceDiscoveryPolicy::EXTENSIONS),
617 no_skills: discovery.disables(ResourceDiscoveryPolicy::SKILLS),
618 no_prompt_templates: discovery.disables(ResourceDiscoveryPolicy::PROMPT_TEMPLATES),
619 no_themes: discovery.disables(ResourceDiscoveryPolicy::THEMES),
620 no_context_files: discovery.disables(ResourceDiscoveryPolicy::CONTEXT_FILES),
621 system_prompt: options.system_prompt,
622 append_system_prompt: options.append_system_prompt,
623 });
624 loader
625 .reload()
626 .await
627 .map_err(|error| AgentSessionServicesError::ResourceLoader(error.to_string()))?;
628 Ok((loader, discovery))
629}
630
631fn extension_discovery_diagnostics(
632 loader: &DefaultResourceLoader,
633) -> Vec<AgentSessionRuntimeDiagnostic> {
634 loader
635 .get_extensions()
636 .errors
637 .iter()
638 .map(|error| {
639 AgentSessionRuntimeDiagnostic::error(format!(
640 "Extension \"{}\" error: {}",
641 error.path, error.error
642 ))
643 })
644 .collect()
645}
646
647async fn start_extension_phase(
648 loader: &DefaultResourceLoader,
649 discovery: ResourceDiscoveryPolicy,
650 cwd: &Path,
651 model_runtime: &ModelRuntime,
652 project_trusted: bool,
653 diagnostics: &mut Vec<AgentSessionRuntimeDiagnostic>,
654) -> (
655 Option<Arc<HostExtensionRunner>>,
656 BTreeMap<String, ExtensionFlagType>,
657) {
658 if discovery.disables(ResourceDiscoveryPolicy::EXTENSIONS) {
659 return (None, BTreeMap::new());
660 }
661 let paths = loader
662 .get_extensions()
663 .paths
664 .iter()
665 .map(|info| {
666 if info.resolved_path.is_empty() {
667 info.path.clone()
668 } else {
669 info.resolved_path.clone()
670 }
671 })
672 .collect::<Vec<_>>();
673 if paths.is_empty() {
674 return (None, BTreeMap::new());
675 }
676 match HostExtensionRunner::start_with_cwd_and_trust(
677 paths,
678 cwd.to_string_lossy().into_owned(),
679 project_trusted,
680 )
681 .await
682 {
683 Ok(runner) => {
684 for (path, message) in runner.load_errors() {
685 diagnostics.push(AgentSessionRuntimeDiagnostic::error(format!(
686 "Extension \"{path}\" error: {message}"
687 )));
688 }
689 for (path, outcome) in runner.register_providers_on(model_runtime) {
690 if let Err(error) = outcome {
691 diagnostics.push(AgentSessionRuntimeDiagnostic::error(format!(
692 "Extension \"{path}\" error: {error}"
693 )));
694 }
695 }
696 let flags = runner.registered_flag_types();
697 (Some(runner), flags)
698 }
699 Err(error) => {
700 diagnostics.push(AgentSessionRuntimeDiagnostic::error(format!(
701 "Extension host failed to start: {error}"
702 )));
703 (None, BTreeMap::new())
704 }
705 }
706}
707
708async fn apply_flags_to_runner(
709 runner: &HostExtensionRunner,
710 applied_flags: &BTreeMap<String, ExtensionFlagValue>,
711) -> Result<(), pi_ext::client::HostClientError> {
712 let values = applied_flags
713 .iter()
714 .map(|(name, value)| {
715 let value = match value {
716 ExtensionFlagValue::Bool(value) => FlagValueWire::Boolean(*value),
717 ExtensionFlagValue::Str(value) => FlagValueWire::String(value.clone()),
718 };
719 (name.clone(), value)
720 })
721 .collect();
722 runner.apply_flag_values(&values).await
723}
724
725#[must_use]
731pub fn apply_extension_flag_values(
732 extension_flag_values: BTreeMap<String, ExtensionFlagValue>,
733 registered_flags: &BTreeMap<String, ExtensionFlagType>,
734) -> (
735 Vec<AgentSessionRuntimeDiagnostic>,
736 BTreeMap<String, ExtensionFlagValue>,
737) {
738 if extension_flag_values.is_empty() {
739 return (Vec::new(), BTreeMap::new());
740 }
741
742 let mut diagnostics = Vec::new();
743 let mut applied = BTreeMap::new();
744 let mut unknown_flags = Vec::new();
745
746 for (name, value) in extension_flag_values {
747 let Some(flag_type) = registered_flags.get(&name) else {
748 unknown_flags.push(name);
749 continue;
750 };
751 match flag_type {
752 ExtensionFlagType::Boolean => {
753 applied.insert(name, ExtensionFlagValue::Bool(true));
754 }
755 ExtensionFlagType::String => match value {
756 ExtensionFlagValue::Str(text) => {
757 applied.insert(name, ExtensionFlagValue::Str(text));
758 }
759 ExtensionFlagValue::Bool(_) => {
760 diagnostics.push(AgentSessionRuntimeDiagnostic::error(format!(
761 "Extension flag \"--{name}\" requires a value"
762 )));
763 }
764 },
765 }
766 }
767
768 if !unknown_flags.is_empty() {
769 let plural = if unknown_flags.len() == 1 { "" } else { "s" };
770 let list = unknown_flags
771 .iter()
772 .map(|name| format!("--{name}"))
773 .collect::<Vec<_>>()
774 .join(", ");
775 diagnostics.push(AgentSessionRuntimeDiagnostic::error(format!(
776 "Unknown option{plural}: {list}"
777 )));
778 }
779
780 (diagnostics, applied)
781}
782
783pub async fn find_initial_model(options: FindInitialModelOptions<'_>) -> InitialModelResult {
790 if let Some(model) = options.cli_model {
791 return InitialModelResult {
792 model: Some(model.clone()),
793 thinking_level: DEFAULT_THINKING_LEVEL,
794 fallback_message: None,
795 };
796 }
797
798 if !options.scoped_models.is_empty() && !options.is_continuing {
799 let first = &options.scoped_models[0];
800 return InitialModelResult {
801 model: Some(first.model.clone()),
802 thinking_level: first.thinking_level.unwrap_or(
803 options
804 .default_thinking_level
805 .unwrap_or(DEFAULT_THINKING_LEVEL),
806 ),
807 fallback_message: None,
808 };
809 }
810
811 if let (Some(provider), Some(model_id)) = (options.default_provider, options.default_model_id)
812 && let Some(found) = options.model_runtime.get_model(provider, model_id)
813 && options.model_runtime.has_configured_auth(&found.provider)
814 {
815 return InitialModelResult {
816 model: Some(found),
817 thinking_level: options
818 .default_thinking_level
819 .unwrap_or(DEFAULT_THINKING_LEVEL),
820 fallback_message: None,
821 };
822 }
823
824 let available = options
825 .model_runtime
826 .get_available(None)
827 .await
828 .unwrap_or_default();
829 if let Some(model) = pick_default_available(&available) {
830 return InitialModelResult {
831 model: Some(model),
832 thinking_level: DEFAULT_THINKING_LEVEL,
833 fallback_message: None,
834 };
835 }
836
837 InitialModelResult {
838 model: None,
839 thinking_level: DEFAULT_THINKING_LEVEL,
840 fallback_message: None,
841 }
842}
843
844pub struct FindInitialModelOptions<'a> {
846 pub cli_model: Option<&'a Model>,
848 pub scoped_models: &'a [ScopedModel],
850 pub is_continuing: bool,
852 pub default_provider: Option<&'a str>,
854 pub default_model_id: Option<&'a str>,
856 pub default_thinking_level: Option<ModelThinkingLevel>,
858 pub model_runtime: &'a ModelRuntime,
860}
861
862pub async fn restore_model_from_session(
864 saved_provider: &str,
865 saved_model_id: &str,
866 current_model: Option<&Model>,
867 model_runtime: &ModelRuntime,
868) -> (Option<Model>, Option<String>) {
869 let restored = model_runtime.get_model(saved_provider, saved_model_id);
870 let has_configured_auth = restored
871 .as_ref()
872 .is_some_and(|model| model_runtime.has_configured_auth(&model.provider));
873
874 if has_configured_auth && let Some(model) = restored {
875 return (Some(model), None);
876 }
877
878 let reason = if restored.is_none() {
879 "model no longer exists"
880 } else {
881 "no auth configured"
882 };
883
884 if let Some(current) = current_model {
885 return (
886 Some(current.clone()),
887 Some(format!(
888 "Could not restore model {saved_provider}/{saved_model_id} ({reason}). Using {}/{}.",
889 current.provider, current.id
890 )),
891 );
892 }
893
894 let available = model_runtime.get_available(None).await.unwrap_or_default();
895 if let Some(fallback) = pick_default_available(&available) {
896 return (
897 Some(fallback.clone()),
898 Some(format!(
899 "Could not restore model {saved_provider}/{saved_model_id} ({reason}). Using {}/{}.",
900 fallback.provider, fallback.id
901 )),
902 );
903 }
904
905 (None, None)
906}
907
908pub async fn create_agent_session_from_services(
919 options: CreateAgentSessionFromServicesOptions,
920) -> Result<CreateAgentSessionResult, AgentSessionServicesError> {
921 let CreateAgentSessionFromServicesOptions {
922 services,
923 model: explicit_model,
924 thinking_level: explicit_thinking,
925 scoped_models,
926 tools,
927 exclude_tools,
928 no_tools,
929 session_start_event,
930 saved_session_model,
931 has_existing_session,
932 } = options;
933
934 let (model, model_fallback_message) = resolve_session_model(
935 &services,
936 explicit_model,
937 &scoped_models,
938 saved_session_model.as_ref(),
939 has_existing_session,
940 )
941 .await;
942 let mut thinking_level = explicit_thinking.unwrap_or_else(|| {
943 services
944 .settings_manager()
945 .get_default_thinking_level()
946 .unwrap_or(DEFAULT_THINKING_LEVEL)
947 });
948 if model.is_none() {
949 thinking_level = ModelThinkingLevel::Off;
950 }
951 let (initial_active_tool_names, allowed_tool_names) =
952 resolve_session_tools(tools, exclude_tools.as_deref(), no_tools);
953 let diagnostics = services.diagnostics.clone();
954
955 Ok(CreateAgentSessionResult {
956 model,
957 thinking_level,
958 initial_active_tool_names,
959 allowed_tool_names,
960 excluded_tool_names: exclude_tools,
961 scoped_models,
962 model_fallback_message,
963 diagnostics,
964 cwd: services.cwd,
965 agent_dir: services.agent_dir,
966 model_runtime: services.model_runtime,
967 resource_loader: services.resource_loader,
968 extension_runner: services.extension_runner,
969 session_start_event,
970 })
971}
972
973async fn resolve_session_model(
974 services: &AgentSessionServices,
975 explicit_model: Option<Model>,
976 scoped_models: &[ScopedModel],
977 saved_session_model: Option<&(String, String)>,
978 has_existing_session: bool,
979) -> (Option<Model>, Option<String>) {
980 let mut model = explicit_model;
981 let mut fallback = None;
982 if model.is_none()
983 && has_existing_session
984 && let Some((provider, model_id)) = saved_session_model
985 {
986 let restored = services.model_runtime.get_model(provider, model_id);
987 if let Some(found) = restored
988 && services.model_runtime.has_configured_auth(&found.provider)
989 {
990 model = Some(found);
991 } else {
992 fallback = Some(format!("Could not restore model {provider}/{model_id}"));
993 }
994 }
995 if model.is_none() {
996 let selected = find_initial_model(FindInitialModelOptions {
997 cli_model: None,
998 scoped_models,
999 is_continuing: has_existing_session,
1000 default_provider: services
1001 .settings_manager()
1002 .get_default_provider()
1003 .as_deref(),
1004 default_model_id: services.settings_manager().get_default_model().as_deref(),
1005 default_thinking_level: services.settings_manager().get_default_thinking_level(),
1006 model_runtime: &services.model_runtime,
1007 })
1008 .await
1009 .model;
1010 match (selected.as_ref(), fallback.as_mut()) {
1011 (None, _) => fallback = Some(format_no_models_available_message()),
1012 (Some(selected), Some(existing)) => {
1013 let _ = write!(existing, ". Using {}/{}", selected.provider, selected.id);
1014 }
1015 (Some(_), None) => {}
1016 }
1017 model = selected;
1018 }
1019 (model, fallback)
1020}
1021
1022fn resolve_session_tools(
1023 tools: Option<Vec<String>>,
1024 exclude_tools: Option<&[String]>,
1025 no_tools: Option<NoToolsMode>,
1026) -> (Vec<String>, Option<Vec<String>>) {
1027 let allowed = match (&tools, no_tools) {
1028 (Some(tools), _) => Some(tools.clone()),
1029 (None, Some(NoToolsMode::All)) => Some(Vec::new()),
1030 (None, Some(NoToolsMode::Builtin) | None) => None,
1031 };
1032 let excluded = exclude_tools.map(|names| names.iter().cloned().collect::<BTreeSet<_>>());
1033 let active = if let Some(tools) = tools {
1034 tools
1035 .into_iter()
1036 .filter(|name| excluded.as_ref().is_none_or(|set| !set.contains(name)))
1037 .collect()
1038 } else if no_tools.is_some() {
1039 Vec::new()
1040 } else {
1041 ["read", "bash", "edit", "write"]
1042 .into_iter()
1043 .filter(|name| excluded.as_ref().is_none_or(|set| !set.contains(*name)))
1044 .map(str::to_owned)
1045 .collect()
1046 };
1047 (active, allowed)
1048}
1049
1050fn pick_default_available(available: &[Model]) -> Option<Model> {
1051 for (provider, default_id) in default_model_per_provider() {
1052 if let Some(match_model) = available
1053 .iter()
1054 .find(|model| model.provider == *provider && model.id == *default_id)
1055 {
1056 return Some(match_model.clone());
1057 }
1058 let _ = KnownProvider::from_id(provider);
1060 }
1061 available.first().cloned()
1062}
1063
1064#[must_use]
1066pub fn resolve_service_path(path: impl AsRef<Path>) -> PathBuf {
1067 resolve_path(path.as_ref().to_string_lossy().as_ref())
1068}
1069
1070#[cfg(test)]
1071mod tests {
1072 use super::*;
1073 use pi_ai::auth::InMemoryCredentialStore;
1074 use pi_ai::models_store::InMemoryModelsStore;
1075 use pi_ai::types::{ModelCost, ModelInput};
1076 use std::io;
1077
1078 use crate::core::model_runtime::{
1079 CreateModelRuntimeOptions, ModelsJsonConfig, ProviderModelDefinition,
1080 };
1081
1082 type TestResult<T = ()> = Result<T, Box<dyn std::error::Error>>;
1083
1084 fn required<T>(value: Option<T>, context: &'static str) -> io::Result<T> {
1085 value.ok_or_else(|| io::Error::other(context))
1086 }
1087
1088 async fn runtime_with_env_openai() -> TestResult<ModelRuntime> {
1089 let mut env = pi_ai::auth::ProviderEnv::new();
1090 env.insert("OPENAI_API_KEY".to_owned(), "sk-test".to_owned());
1091 Ok(ModelRuntime::create(CreateModelRuntimeOptions {
1092 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
1093 models_store: Some(Arc::new(InMemoryModelsStore::new())),
1094 models_config: Some(ModelsJsonConfig::empty()),
1095 allow_model_network: Some(false),
1096 auth_env: Some(env),
1097 ..CreateModelRuntimeOptions::default()
1098 })
1099 .await?)
1100 }
1101
1102 #[test]
1103 fn auth_guidance_strings_match_typescript() {
1104 let help = get_provider_login_help();
1105 assert!(help.contains("Use /login to log into a provider via OAuth or API key."));
1106 assert!(help.contains("providers.md"));
1107 assert!(help.contains("models.md"));
1108
1109 let no_models = format_no_models_available_message();
1110 assert!(no_models.starts_with("No models available. "));
1111 assert!(no_models.contains(&help));
1112
1113 let no_selected = format_no_model_selected_message();
1114 assert!(no_selected.starts_with("No model selected."));
1115 assert!(no_selected.contains("Then use /model to select a model."));
1116
1117 let no_key = format_no_api_key_found_message("anthropic");
1118 assert_eq!(no_key, format!("No API key found for anthropic.\n\n{help}"));
1119 let unknown = format_no_api_key_found_message("unknown");
1120 assert!(unknown.contains("the selected model"));
1121
1122 let oauth = format_oauth_auth_failed_message("openai-codex");
1123 assert_eq!(
1124 oauth,
1125 "Authentication failed for \"openai-codex\". Credentials may have expired or network is unavailable. Run '/login openai-codex' to re-authenticate."
1126 );
1127 }
1128
1129 #[test]
1130 fn extension_flag_validation_unknown_and_string_required() {
1131 let mut flags = BTreeMap::new();
1132 flags.insert("verbose".to_owned(), ExtensionFlagValue::Bool(true));
1133 flags.insert("mode".to_owned(), ExtensionFlagValue::Bool(true));
1134 flags.insert("unknown".to_owned(), ExtensionFlagValue::Str("x".into()));
1135
1136 let mut registered = BTreeMap::new();
1137 registered.insert("verbose".to_owned(), ExtensionFlagType::Boolean);
1138 registered.insert("mode".to_owned(), ExtensionFlagType::String);
1139
1140 let (diagnostics, applied) = apply_extension_flag_values(flags, ®istered);
1141 assert!(applied.contains_key("verbose"));
1142 assert!(!applied.contains_key("mode"));
1143 assert_eq!(diagnostics.len(), 2);
1144 assert!(
1145 diagnostics
1146 .iter()
1147 .any(|d| d.message == "Extension flag \"--mode\" requires a value")
1148 );
1149 assert!(
1150 diagnostics
1151 .iter()
1152 .any(|d| d.message == "Unknown option: --unknown")
1153 );
1154 }
1155
1156 #[test]
1157 fn extension_flag_validation_plural_unknown_options() {
1158 let mut flags = BTreeMap::new();
1159 flags.insert("a".to_owned(), ExtensionFlagValue::Bool(true));
1160 flags.insert("b".to_owned(), ExtensionFlagValue::Str("1".into()));
1161 let (diagnostics, _) = apply_extension_flag_values(flags, &BTreeMap::new());
1162 assert_eq!(diagnostics.len(), 1);
1163 assert_eq!(diagnostics[0].message, "Unknown options: --a, --b");
1164 }
1165
1166 #[tokio::test]
1167 async fn services_creation_order_registers_pending_providers() -> TestResult {
1168 let dir = tempfile::tempdir()?;
1169 let cwd = dir.path().join("project");
1170 let agent = dir.path().join("agent");
1171 std::fs::create_dir_all(&cwd)?;
1172 std::fs::create_dir_all(&agent)?;
1173
1174 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
1175 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
1176 models_store: Some(Arc::new(InMemoryModelsStore::new())),
1177 models_config: Some(ModelsJsonConfig::empty()),
1178 allow_model_network: Some(false),
1179 ..CreateModelRuntimeOptions::default()
1180 })
1181 .await?;
1182
1183 let services = create_agent_session_services(CreateAgentSessionServicesOptions {
1184 cwd: cwd.clone(),
1185 agent_dir: Some(agent.clone()),
1186 model_runtime: Some(runtime),
1187 pending_provider_registrations: vec![PendingProviderRegistration {
1188 name: "acme".to_owned(),
1189 config: ProviderConfigInput {
1190 base_url: Some("https://acme.test/v1".into()),
1191 api: Some("openai-completions".into()),
1192 api_key: Some("sk-acme".into()),
1193 models: Some(vec![ProviderModelDefinition {
1194 id: "acme-1".into(),
1195 name: Some("Acme 1".into()),
1196 api: Some("openai-completions".into()),
1197 base_url: Some("https://acme.test/v1".into()),
1198 reasoning: false,
1199 thinking_level_map: None,
1200 input: Some(vec![ModelInput::Text]),
1201 cost: Some(ModelCost::default()),
1202 context_window: Some(8_000),
1203 max_tokens: Some(1_024),
1204 headers: None,
1205 compat: None,
1206 }]),
1207 ..ProviderConfigInput::default()
1208 },
1209 extension_path: "/ext/acme.ts".into(),
1210 }],
1211 registered_extension_flags: BTreeMap::new(),
1212 extension_flag_values: None,
1213 settings_manager: None,
1214 resource_loader_options: Some(ResourceLoaderServiceOptions {
1215 no_extensions: true,
1216 no_skills: true,
1217 no_prompt_templates: true,
1218 no_themes: true,
1219 no_context_files: true,
1220 ..ResourceLoaderServiceOptions::default()
1221 }),
1222 })
1223 .await?;
1224
1225 assert!(services.model_runtime.get_model("acme", "acme-1").is_some());
1226 assert!(services.model_runtime.has_configured_auth("acme"));
1227 assert!(services.diagnostics.is_empty());
1228 Ok(())
1229 }
1230
1231 #[tokio::test]
1232 async fn pending_provider_failure_becomes_diagnostic() -> TestResult {
1233 let dir = tempfile::tempdir()?;
1234 let cwd = dir.path().join("project");
1235 let agent = dir.path().join("agent");
1236 std::fs::create_dir_all(&cwd)?;
1237 std::fs::create_dir_all(&agent)?;
1238
1239 let runtime = ModelRuntime::create_in_memory().await?;
1240 let services = create_agent_session_services(CreateAgentSessionServicesOptions {
1241 cwd,
1242 agent_dir: Some(agent),
1243 model_runtime: Some(runtime),
1244 pending_provider_registrations: vec![PendingProviderRegistration {
1245 name: "broken".into(),
1246 config: ProviderConfigInput {
1247 models: Some(vec![ProviderModelDefinition {
1248 id: "m".into(),
1249 name: None,
1250 api: None,
1251 base_url: None,
1252 reasoning: false,
1253 thinking_level_map: None,
1254 input: None,
1255 cost: None,
1256 context_window: None,
1257 max_tokens: None,
1258 headers: None,
1259 compat: None,
1260 }]),
1261 ..ProviderConfigInput::default()
1262 },
1263 extension_path: "/ext/broken.ts".into(),
1264 }],
1265 resource_loader_options: Some(ResourceLoaderServiceOptions {
1266 no_extensions: true,
1267 no_skills: true,
1268 no_prompt_templates: true,
1269 no_themes: true,
1270 no_context_files: true,
1271 ..ResourceLoaderServiceOptions::default()
1272 }),
1273 ..CreateAgentSessionServicesOptions::default()
1274 })
1275 .await?;
1276
1277 assert_eq!(services.diagnostics.len(), 1);
1278 assert!(
1279 services.diagnostics[0]
1280 .message
1281 .starts_with("Extension \"/ext/broken.ts\" error:")
1282 );
1283 Ok(())
1284 }
1285
1286 #[tokio::test]
1287 async fn find_initial_model_priority_settings_then_available() -> TestResult {
1288 let runtime = runtime_with_env_openai().await?;
1289 let openai_default = required(
1290 runtime.get_model("openai", "gpt-5.5"),
1291 "OpenAI default must exist in the built-in catalog",
1292 )?;
1293
1294 let result = find_initial_model(FindInitialModelOptions {
1296 cli_model: None,
1297 scoped_models: &[],
1298 is_continuing: false,
1299 default_provider: Some("openai"),
1300 default_model_id: Some("gpt-5.5"),
1301 default_thinking_level: Some(ModelThinkingLevel::High),
1302 model_runtime: &runtime,
1303 })
1304 .await;
1305 assert_eq!(
1306 result.model.as_ref().map(|m| m.id.as_str()),
1307 Some("gpt-5.5")
1308 );
1309 assert_eq!(result.thinking_level, ModelThinkingLevel::High);
1310
1311 let cli = openai_default.clone();
1313 let result = find_initial_model(FindInitialModelOptions {
1314 cli_model: Some(&cli),
1315 scoped_models: &[],
1316 is_continuing: false,
1317 default_provider: Some("openai"),
1318 default_model_id: Some("other"),
1319 default_thinking_level: None,
1320 model_runtime: &runtime,
1321 })
1322 .await;
1323 assert_eq!(
1324 result.model.as_ref().map(|m| m.id.as_str()),
1325 Some("gpt-5.5")
1326 );
1327 assert_eq!(result.thinking_level, DEFAULT_THINKING_LEVEL);
1328 Ok(())
1329 }
1330
1331 #[tokio::test]
1332 async fn find_initial_model_scoped_skipped_when_continuing() -> TestResult {
1333 let runtime = runtime_with_env_openai().await?;
1334 let model = required(
1335 runtime.get_model("openai", "gpt-5.5"),
1336 "OpenAI test model must exist in the built-in catalog",
1337 )?;
1338 let scoped = vec![ScopedModel {
1339 model: model.clone(),
1340 thinking_level: Some(ModelThinkingLevel::Low),
1341 }];
1342
1343 let continuing = find_initial_model(FindInitialModelOptions {
1344 cli_model: None,
1345 scoped_models: &scoped,
1346 is_continuing: true,
1347 default_provider: None,
1348 default_model_id: None,
1349 default_thinking_level: None,
1350 model_runtime: &runtime,
1351 })
1352 .await;
1353 assert_eq!(
1355 continuing.model.as_ref().map(|m| m.id.as_str()),
1356 Some("gpt-5.5")
1357 );
1358 assert_eq!(continuing.thinking_level, DEFAULT_THINKING_LEVEL);
1359
1360 let fresh = find_initial_model(FindInitialModelOptions {
1361 cli_model: None,
1362 scoped_models: &scoped,
1363 is_continuing: false,
1364 default_provider: None,
1365 default_model_id: None,
1366 default_thinking_level: None,
1367 model_runtime: &runtime,
1368 })
1369 .await;
1370 assert_eq!(fresh.thinking_level, ModelThinkingLevel::Low);
1371 Ok(())
1372 }
1373
1374 #[tokio::test]
1375 async fn restore_model_fallback_message() -> TestResult {
1376 let runtime = runtime_with_env_openai().await?;
1377 let current = required(
1378 runtime.get_model("openai", "gpt-5.5"),
1379 "OpenAI fallback model must exist in the built-in catalog",
1380 )?;
1381 let (model, message) =
1382 restore_model_from_session("missing", "gone", Some(¤t), &runtime).await;
1383 assert_eq!(model.as_ref().map(|m| m.id.as_str()), Some("gpt-5.5"));
1384 assert_eq!(
1385 message.as_deref(),
1386 Some(
1387 "Could not restore model missing/gone (model no longer exists). Using openai/gpt-5.5."
1388 )
1389 );
1390 Ok(())
1391 }
1392
1393 #[tokio::test]
1394 async fn create_session_from_services_tool_resolution_and_fallback() -> TestResult {
1395 let dir = tempfile::tempdir()?;
1396 let cwd = dir.path().join("project");
1397 let agent = dir.path().join("agent");
1398 std::fs::create_dir_all(&cwd)?;
1399 std::fs::create_dir_all(&agent)?;
1400
1401 let runtime = ModelRuntime::create_in_memory().await?;
1402 let services = create_agent_session_services(CreateAgentSessionServicesOptions {
1403 cwd,
1404 agent_dir: Some(agent),
1405 model_runtime: Some(runtime),
1406 resource_loader_options: Some(ResourceLoaderServiceOptions {
1407 no_extensions: true,
1408 no_skills: true,
1409 no_prompt_templates: true,
1410 no_themes: true,
1411 no_context_files: true,
1412 ..ResourceLoaderServiceOptions::default()
1413 }),
1414 ..CreateAgentSessionServicesOptions::default()
1415 })
1416 .await?;
1417
1418 let result = create_agent_session_from_services(CreateAgentSessionFromServicesOptions {
1419 services,
1420 model: None,
1421 thinking_level: None,
1422 scoped_models: Vec::new(),
1423 tools: None,
1424 exclude_tools: Some(vec!["bash".into()]),
1425 no_tools: None,
1426 session_start_event: None,
1427 saved_session_model: None,
1428 has_existing_session: false,
1429 })
1430 .await?;
1431
1432 assert!(result.model.is_none());
1433 assert_eq!(
1434 result.model_fallback_message.as_deref(),
1435 Some(format_no_models_available_message().as_str())
1436 );
1437 assert_eq!(
1438 result.initial_active_tool_names,
1439 vec!["read".to_owned(), "edit".to_owned(), "write".to_owned()]
1440 );
1441 assert_eq!(result.thinking_level, ModelThinkingLevel::Off);
1442 assert!(result.allowed_tool_names.is_none());
1443 Ok(())
1444 }
1445
1446 #[tokio::test]
1447 async fn from_services_forwards_session_start_event() -> TestResult {
1448 let dir = tempfile::tempdir()?;
1449 let cwd = dir.path().join("project");
1450 let agent = dir.path().join("agent");
1451 std::fs::create_dir_all(&cwd)?;
1452 std::fs::create_dir_all(&agent)?;
1453
1454 let runtime = ModelRuntime::create_in_memory().await?;
1455 let services = create_agent_session_services(CreateAgentSessionServicesOptions {
1456 cwd,
1457 agent_dir: Some(agent),
1458 model_runtime: Some(runtime),
1459 resource_loader_options: Some(ResourceLoaderServiceOptions {
1460 no_extensions: true,
1461 no_skills: true,
1462 no_prompt_templates: true,
1463 no_themes: true,
1464 no_context_files: true,
1465 ..ResourceLoaderServiceOptions::default()
1466 }),
1467 ..CreateAgentSessionServicesOptions::default()
1468 })
1469 .await?;
1470
1471 let event = crate::core::agent_session::SessionStartEvent {
1472 reason: crate::core::agent_session::SessionStartReason::New,
1473 previous_session_file: Some("prev.jsonl".to_owned()),
1474 };
1475 let result = create_agent_session_from_services(CreateAgentSessionFromServicesOptions {
1476 services,
1477 model: None,
1478 thinking_level: None,
1479 scoped_models: Vec::new(),
1480 tools: None,
1481 exclude_tools: None,
1482 no_tools: None,
1483 session_start_event: Some(event.clone()),
1484 saved_session_model: None,
1485 has_existing_session: false,
1486 })
1487 .await?;
1488
1489 assert_eq!(
1490 result.session_start_event,
1491 Some(event),
1492 "replacement session-start metadata must survive services resolution"
1493 );
1494 Ok(())
1495 }
1496
1497 #[tokio::test]
1498 async fn no_tools_all_sets_empty_allowlist_builtin_leaves_none() -> TestResult {
1499 let dir = tempfile::tempdir()?;
1500 let cwd = dir.path().join("project");
1501 let agent = dir.path().join("agent");
1502 std::fs::create_dir_all(&cwd)?;
1503 std::fs::create_dir_all(&agent)?;
1504
1505 let runtime = ModelRuntime::create_in_memory().await?;
1506 let services = create_agent_session_services(CreateAgentSessionServicesOptions {
1507 cwd: cwd.clone(),
1508 agent_dir: Some(agent.clone()),
1509 model_runtime: Some(runtime.clone()),
1510 resource_loader_options: Some(ResourceLoaderServiceOptions {
1511 no_extensions: true,
1512 no_skills: true,
1513 no_prompt_templates: true,
1514 no_themes: true,
1515 no_context_files: true,
1516 ..ResourceLoaderServiceOptions::default()
1517 }),
1518 ..CreateAgentSessionServicesOptions::default()
1519 })
1520 .await?;
1521
1522 let all = create_agent_session_from_services(CreateAgentSessionFromServicesOptions {
1523 services,
1524 model: None,
1525 thinking_level: None,
1526 scoped_models: Vec::new(),
1527 tools: None,
1528 exclude_tools: None,
1529 no_tools: Some(NoToolsMode::All),
1530 session_start_event: None,
1531 saved_session_model: None,
1532 has_existing_session: false,
1533 })
1534 .await?;
1535 assert_eq!(all.allowed_tool_names, Some(Vec::new()));
1536 assert!(all.initial_active_tool_names.is_empty());
1537
1538 let runtime = ModelRuntime::create_in_memory().await?;
1539 let services = create_agent_session_services(CreateAgentSessionServicesOptions {
1540 cwd,
1541 agent_dir: Some(agent),
1542 model_runtime: Some(runtime),
1543 resource_loader_options: Some(ResourceLoaderServiceOptions {
1544 no_extensions: true,
1545 no_skills: true,
1546 no_prompt_templates: true,
1547 no_themes: true,
1548 no_context_files: true,
1549 ..ResourceLoaderServiceOptions::default()
1550 }),
1551 ..CreateAgentSessionServicesOptions::default()
1552 })
1553 .await?;
1554
1555 let builtin = create_agent_session_from_services(CreateAgentSessionFromServicesOptions {
1556 services,
1557 model: None,
1558 thinking_level: None,
1559 scoped_models: Vec::new(),
1560 tools: None,
1561 exclude_tools: None,
1562 no_tools: Some(NoToolsMode::Builtin),
1563 session_start_event: None,
1564 saved_session_model: None,
1565 has_existing_session: false,
1566 })
1567 .await?;
1568 assert!(builtin.allowed_tool_names.is_none());
1569 assert!(builtin.initial_active_tool_names.is_empty());
1570 Ok(())
1571 }
1572 #[tokio::test]
1573 async fn trust_override_precedes_project_settings_load() -> TestResult {
1574 for (trust_override, trusted) in [(None, false), (Some(false), false), (Some(true), true)] {
1575 let dir = tempfile::tempdir()?;
1576 let cwd = dir.path().join("project");
1577 let agent = dir.path().join("agent");
1578 std::fs::create_dir_all(cwd.join(".pi"))?;
1579 std::fs::create_dir_all(&agent)?;
1580 std::fs::write(agent.join("settings.json"), r#"{"theme":"global"}"#)?;
1581 std::fs::write(
1582 cwd.join(".pi").join("settings.json"),
1583 r#"{"theme":"project"}"#,
1584 )?;
1585
1586 let services = create_agent_session_services_with_trust(
1587 CreateAgentSessionServicesOptions {
1588 cwd,
1589 agent_dir: Some(agent),
1590 model_runtime: Some(ModelRuntime::create_in_memory().await?),
1591 resource_loader_options: Some(ResourceLoaderServiceOptions {
1592 no_extensions: true,
1593 no_skills: true,
1594 no_prompt_templates: true,
1595 no_themes: true,
1596 no_context_files: true,
1597 ..Default::default()
1598 }),
1599 ..Default::default()
1600 },
1601 trust_override,
1602 )
1603 .await?;
1604
1605 assert_eq!(services.settings_manager().is_project_trusted(), trusted);
1606 assert_eq!(
1607 services.settings_manager().get_theme().as_deref(),
1608 Some(if trusted { "project" } else { "global" })
1609 );
1610 }
1611 Ok(())
1612 }
1613}