1use std::collections::HashMap;
2use std::path::{Path, PathBuf};
3use std::sync::Arc;
4use std::time::{Duration, Instant};
5
6use parking_lot::Mutex as ParkingLotMutex;
7use serde_json::Value;
8use tokio::sync::oneshot;
9use tokio::task::JoinHandle;
10use tokio_util::sync::CancellationToken;
11use tracing::{Instrument, warn};
12
13use crate::canvas::CanvasHandler;
14use crate::generated::api_types::{
15 LogRequest, ModelSwitchToRequest, OpenCanvasInstance, PermissionDecisionRequest,
16 RegisterEventInterestParams, ToolsGetCurrentMetadataResult, rpc_methods,
17};
18use crate::generated::session_events::{
19 CommandExecuteData, ElicitationRequestedData, ExternalToolRequestedData, McpOauthRequiredData,
20 SessionCanvasClosedData, SessionErrorData, SessionEventType,
21};
22use crate::handler::{
23 AutoModeSwitchHandler, AutoModeSwitchResponse, ElicitationHandler, ExitPlanModeHandler,
24 McpAuthHandler, McpAuthRequest, McpAuthResult, PermissionHandler, PermissionResult,
25 UserInputHandler, UserInputResponse,
26};
27use crate::hooks::SessionHooks;
28use crate::provider_token::BearerTokenProvider;
29use crate::session_fs::SessionFsProvider;
30use crate::trace_context::inject_trace_context;
31use crate::transforms::SystemMessageTransform;
32use crate::types::{
33 CommandContext, CommandDefinition, CommandHandler, CreateSessionResult, ElicitationRequest,
34 ElicitationResult, ExitPlanModeData, GetMessagesResponse, MessageOptions,
35 PermissionRequestData, RequestId, ResumeSessionConfig, ResumeSessionResult, SectionOverride,
36 SessionCapabilities, SessionConfig, SessionEvent, SessionId, SetModelOptions,
37 SystemMessageConfig, ToolInvocation, ToolResult, ToolResultExpanded, TraceContext,
38 UiInputOptions, ensure_attachment_display_names,
39};
40use crate::{
41 Client, Error, ErrorKind, JsonRpcResponse, SessionErrorKind, SessionEventNotification,
42 error_codes,
43};
44
45const TOOL_SEARCH_TOOL_NAME: &str = "tool_search_tool";
49
50#[derive(Clone)]
58pub(crate) struct SessionHandlers {
59 pub permission: Option<Arc<dyn PermissionHandler>>,
60 pub managed_settings_enabled: bool,
61 pub elicitation: Option<Arc<dyn ElicitationHandler>>,
62 pub mcp_auth: Option<Arc<dyn McpAuthHandler>>,
63 pub user_input: Option<Arc<dyn UserInputHandler>>,
64 pub exit_plan_mode: Option<Arc<dyn ExitPlanModeHandler>>,
65 pub auto_mode_switch: Option<Arc<dyn AutoModeSwitchHandler>>,
66 pub tools: Arc<HashMap<String, Arc<dyn crate::tool::ToolHandler>>>,
67}
68
69fn has_managed_settings(
70 enable_managed_settings: Option<bool>,
71 managed_settings: Option<&crate::types::ManagedSettings>,
72) -> bool {
73 enable_managed_settings == Some(true) || managed_settings.is_some()
74}
75
76struct IdleWaiter {
78 tx: oneshot::Sender<Result<Option<SessionEvent>, Error>>,
79 last_assistant_message: Option<SessionEvent>,
80 started_at: Instant,
81 first_assistant_message_seen: bool,
82}
83
84struct WaiterGuard {
96 slot: Arc<ParkingLotMutex<Option<IdleWaiter>>>,
97}
98
99impl Drop for WaiterGuard {
100 fn drop(&mut self) {
101 self.slot.lock().take();
102 }
103}
104
105struct PendingSessionRegistration {
106 client: Client,
107 session_id: SessionId,
108 shutdown: CancellationToken,
109 disarmed: bool,
110}
111
112impl PendingSessionRegistration {
113 fn new(client: Client, session_id: SessionId, shutdown: CancellationToken) -> Self {
114 Self {
115 client,
116 session_id,
117 shutdown,
118 disarmed: false,
119 }
120 }
121
122 async fn cleanup(mut self, event_loop: JoinHandle<()>) {
123 self.shutdown.cancel();
124 let _ = event_loop.await;
125 self.client.unregister_session(&self.session_id);
126 self.disarmed = true;
127 }
128
129 fn disarm(&mut self) {
130 self.disarmed = true;
131 }
132}
133
134impl Drop for PendingSessionRegistration {
135 fn drop(&mut self) {
136 if !self.disarmed {
137 self.shutdown.cancel();
138 self.client.unregister_session(&self.session_id);
139 }
140 }
141}
142
143pub struct Session {
156 id: SessionId,
157 cwd: PathBuf,
158 workspace_path: Option<PathBuf>,
159 remote_url: Option<String>,
160 client: Client,
161 event_loop: ParkingLotMutex<Option<JoinHandle<()>>>,
166 shutdown: CancellationToken,
180 idle_waiter: Arc<ParkingLotMutex<Option<IdleWaiter>>>,
187 capabilities: Arc<parking_lot::RwLock<SessionCapabilities>>,
189 open_canvases: Arc<parking_lot::RwLock<Vec<OpenCanvasInstance>>>,
191 event_tx: tokio::sync::broadcast::Sender<SessionEvent>,
193}
194
195impl Session {
196 pub fn id(&self) -> &SessionId {
198 &self.id
199 }
200
201 pub fn cwd(&self) -> &PathBuf {
203 &self.cwd
204 }
205
206 pub fn workspace_path(&self) -> Option<&Path> {
208 self.workspace_path.as_deref()
209 }
210
211 pub fn remote_url(&self) -> Option<&str> {
213 self.remote_url.as_deref()
214 }
215
216 pub fn capabilities(&self) -> SessionCapabilities {
221 self.capabilities.read().clone()
222 }
223
224 pub fn open_canvases(&self) -> Vec<OpenCanvasInstance> {
227 self.open_canvases.read().clone()
228 }
229
230 pub fn cancellation_token(&self) -> CancellationToken {
257 self.shutdown.child_token()
258 }
259
260 pub fn subscribe(&self) -> crate::subscription::EventSubscription {
299 crate::subscription::EventSubscription::new(self.event_tx.subscribe())
300 }
301
302 pub fn client(&self) -> &Client {
304 &self.client
305 }
306
307 pub fn rpc(&self) -> crate::generated::rpc::SessionRpc<'_> {
318 crate::generated::rpc::SessionRpc { session: self }
319 }
320
321 pub async fn stop_event_loop(&self) {
329 self.shutdown.cancel();
330 let handle = self.event_loop.lock().take();
331 if let Some(handle) = handle {
332 let _ = handle.await;
333 }
334 if let Some(waiter) = self.idle_waiter.lock().take() {
336 let _ = waiter.tx.send(Err(
337 ErrorKind::Session(SessionErrorKind::EventLoopClosed).into()
338 ));
339 }
340 }
341
342 pub async fn send(&self, opts: impl Into<MessageOptions>) -> Result<String, Error> {
367 if self.idle_waiter.lock().is_some() {
368 return Err(ErrorKind::Session(SessionErrorKind::SendWhileWaiting).into());
369 }
370 self.send_inner(opts.into()).await
371 }
372
373 async fn send_inner(&self, opts: MessageOptions) -> Result<String, Error> {
374 let mut params = serde_json::json!({
375 "sessionId": self.id,
376 "prompt": opts.prompt,
377 });
378 if let Some(m) = opts.mode {
379 params["mode"] = serde_json::to_value(m)?;
380 }
381 if let Some(am) = opts.agent_mode {
382 params["agentMode"] = serde_json::to_value(am)?;
383 }
384 if let Some(mut a) = opts.attachments {
385 ensure_attachment_display_names(&mut a);
386 params["attachments"] = serde_json::to_value(a)?;
387 }
388 if let Some(headers) = opts.request_headers
389 && !headers.is_empty()
390 {
391 params["requestHeaders"] = serde_json::to_value(headers)?;
392 }
393 if let Some(display_prompt) = opts.display_prompt {
394 params["displayPrompt"] = serde_json::to_value(display_prompt)?;
395 }
396 let trace_ctx = if opts.traceparent.is_some() || opts.tracestate.is_some() {
397 TraceContext {
398 traceparent: opts.traceparent,
399 tracestate: opts.tracestate,
400 }
401 } else {
402 self.client.resolve_trace_context().await
403 };
404 inject_trace_context(&mut params, &trace_ctx);
405 let rpc_start = Instant::now();
406 let result = self.client.call("session.send", Some(params)).await?;
407 let message_id = result
408 .get("messageId")
409 .and_then(|v| v.as_str())
410 .map(|s| s.to_string())
411 .unwrap_or_default();
412 tracing::debug!(
413 elapsed_ms = rpc_start.elapsed().as_millis(),
414 session_id = %self.id,
415 message_id = %message_id,
416 "Session::send completed successfully"
417 );
418 Ok(message_id)
419 }
420
421 pub async fn send_and_wait(
441 &self,
442 opts: impl Into<MessageOptions>,
443 ) -> Result<Option<SessionEvent>, Error> {
444 let total_start = Instant::now();
445 let opts = opts.into();
446 let timeout_duration = opts.wait_timeout.unwrap_or(Duration::from_secs(60));
447 let (tx, rx) = oneshot::channel();
448
449 {
450 let mut guard = self.idle_waiter.lock();
451 if guard.is_some() {
452 return Err(ErrorKind::Session(SessionErrorKind::SendWhileWaiting).into());
453 }
454 *guard = Some(IdleWaiter {
455 tx,
456 last_assistant_message: None,
457 started_at: total_start,
458 first_assistant_message_seen: false,
459 });
460 }
461
462 let _waiter_guard = WaiterGuard {
467 slot: self.idle_waiter.clone(),
468 };
469
470 let result = tokio::time::timeout(timeout_duration, async {
471 self.send_inner(opts).await?;
472 match rx.await {
473 Ok(result) => result,
474 Err(_) => Err(ErrorKind::Session(SessionErrorKind::EventLoopClosed).into()),
475 }
476 })
477 .await;
478
479 match result {
480 Ok(inner) => {
481 tracing::debug!(
482 elapsed_ms = total_start.elapsed().as_millis(),
483 session_id = %self.id,
484 completed_by = if inner.is_ok() { "idle" } else { "error" },
485 "Session::send_and_wait complete"
486 );
487 inner
488 }
489 Err(_) => {
490 tracing::warn!(
491 elapsed_ms = total_start.elapsed().as_millis(),
492 session_id = %self.id,
493 completed_by = "timeout",
494 "Session::send_and_wait failed"
495 );
496 Err(ErrorKind::Session(SessionErrorKind::Timeout(timeout_duration)).into())
497 }
498 }
499 }
500
501 pub async fn get_events(&self) -> Result<Vec<SessionEvent>, Error> {
503 let result = self
504 .client
505 .call(
506 "session.getMessages",
507 Some(serde_json::json!({ "sessionId": self.id })),
508 )
509 .await?;
510 let response: GetMessagesResponse = serde_json::from_value(result)?;
511 Ok(response.events)
512 }
513
514 #[deprecated(since = "0.1.0", note = "Use `get_events()` instead")]
516 pub async fn get_messages(&self) -> Result<Vec<SessionEvent>, Error> {
517 self.get_events().await
518 }
519
520 pub async fn abort(&self) -> Result<(), Error> {
528 self.client
529 .call(
530 "session.abort",
531 Some(serde_json::json!({ "sessionId": self.id })),
532 )
533 .await?;
534 Ok(())
535 }
536
537 pub async fn set_model(&self, model: &str, opts: Option<SetModelOptions>) -> Result<(), Error> {
541 let opts = opts.unwrap_or_default();
542 let request = ModelSwitchToRequest {
543 compaction_decision: None,
544 context_tier: opts.context_tier,
545 defer_if_model_change_queued: None,
546 model_capabilities: opts.model_capabilities,
547 model_change_scope: None,
548 model_id: model.to_string(),
549 picker_persistence: None,
550 reasoning_effort: opts.reasoning_effort,
551 reasoning_summary: opts.reasoning_summary,
552 repo_scope: None,
553 require_available: None,
554 run_compaction_preflight: None,
555 source: None,
556 verbosity: None,
557 };
558 self.rpc().model().switch_to(request).await?;
559 Ok(())
560 }
561
562 pub async fn disconnect(&self) -> Result<(), Error> {
579 self.client
580 .call(
581 "session.destroy",
582 Some(serde_json::json!({ "sessionId": self.id })),
583 )
584 .await?;
585 self.stop_event_loop().await;
586 self.client.unregister_session(&self.id);
587 Ok(())
588 }
589
590 #[deprecated(since = "0.1.0", note = "Use `disconnect()` instead")]
595 pub async fn destroy(&self) -> Result<(), Error> {
596 self.disconnect().await
597 }
598
599 pub async fn log(
603 &self,
604 message: &str,
605 opts: Option<crate::types::LogOptions>,
606 ) -> Result<(), Error> {
607 let opts = opts.unwrap_or_default();
608 let level = match opts.level {
609 Some(level) => Some(serde_json::from_value(serde_json::to_value(level)?)?),
610 None => None,
611 };
612 let request = LogRequest {
613 message: message.to_string(),
614 level,
615 ephemeral: opts.ephemeral,
616 r#type: None,
617 tip: None,
618 url: None,
619 };
620 self.rpc().log(request).await?;
621 Ok(())
622 }
623
624 pub fn ui(&self) -> SessionUi<'_> {
630 SessionUi { session: self }
631 }
632
633 fn assert_elicitation(&self) -> Result<(), Error> {
635 if self
636 .capabilities
637 .read()
638 .ui
639 .as_ref()
640 .and_then(|u| u.elicitation)
641 != Some(true)
642 {
643 return Err(ErrorKind::Session(SessionErrorKind::ElicitationNotSupported).into());
644 }
645 Ok(())
646 }
647}
648
649impl Drop for Session {
650 fn drop(&mut self) {
651 self.shutdown.cancel();
663 self.client.unregister_session(&self.id);
664 }
665}
666
667pub struct SessionUi<'a> {
674 session: &'a Session,
675}
676
677impl<'a> SessionUi<'a> {
678 pub async fn elicitation(
686 &self,
687 message: &str,
688 schema: Value,
689 ) -> Result<ElicitationResult, Error> {
690 self.session.assert_elicitation()?;
691 let result = self
692 .session
693 .client
694 .call(
695 "session.ui.elicitation",
696 Some(serde_json::json!({
697 "sessionId": self.session.id,
698 "message": message,
699 "requestedSchema": schema,
700 })),
701 )
702 .await?;
703 let elicitation: ElicitationResult = serde_json::from_value(result)?;
704 Ok(elicitation)
705 }
706
707 pub async fn confirm(&self, message: &str) -> Result<bool, Error> {
711 self.session.assert_elicitation()?;
712 let schema = serde_json::json!({
713 "type": "object",
714 "properties": {
715 "confirmed": {
716 "type": "boolean",
717 "default": true,
718 }
719 },
720 "required": ["confirmed"]
721 });
722 let result = self.elicitation(message, schema).await?;
723 Ok(result.action == "accept"
724 && result
725 .content
726 .and_then(|c| c.get("confirmed").and_then(|v| v.as_bool()))
727 == Some(true))
728 }
729
730 pub async fn select(&self, message: &str, options: &[&str]) -> Result<Option<String>, Error> {
734 self.session.assert_elicitation()?;
735 let schema = serde_json::json!({
736 "type": "object",
737 "properties": {
738 "selection": {
739 "type": "string",
740 "enum": options,
741 }
742 },
743 "required": ["selection"]
744 });
745 let result = self.elicitation(message, schema).await?;
746 if result.action != "accept" {
747 return Ok(None);
748 }
749 let selection = result.content.and_then(|c| {
750 c.get("selection")
751 .and_then(|v| v.as_str())
752 .map(String::from)
753 });
754 Ok(selection)
755 }
756
757 pub async fn input(
762 &self,
763 message: &str,
764 options: Option<&UiInputOptions<'_>>,
765 ) -> Result<Option<String>, Error> {
766 self.session.assert_elicitation()?;
767 let mut field = serde_json::json!({ "type": "string" });
768 if let Some(opts) = options {
769 if let Some(title) = opts.title {
770 field["title"] = Value::String(title.to_string());
771 }
772 if let Some(desc) = opts.description {
773 field["description"] = Value::String(desc.to_string());
774 }
775 if let Some(min) = opts.min_length {
776 field["minLength"] = Value::Number(min.into());
777 }
778 if let Some(max) = opts.max_length {
779 field["maxLength"] = Value::Number(max.into());
780 }
781 if let Some(fmt) = &opts.format {
782 field["format"] = Value::String(fmt.as_str().to_string());
783 }
784 if let Some(default) = opts.default {
785 field["default"] = Value::String(default.to_string());
786 }
787 }
788 let schema = serde_json::json!({
789 "type": "object",
790 "properties": { "value": field },
791 "required": ["value"]
792 });
793 let result = self.elicitation(message, schema).await?;
794 if result.action != "accept" {
795 return Ok(None);
796 }
797 let value = result
798 .content
799 .and_then(|c| c.get("value").and_then(|v| v.as_str()).map(String::from));
800 Ok(value)
801 }
802}
803
804impl Client {
805 pub async fn create_session(&self, mut config: SessionConfig) -> Result<Session, Error> {
827 let total_start = Instant::now();
828 let caller_session_id = config.session_id.clone();
836 let use_server_generated_id = config.cloud.is_some() && caller_session_id.is_none();
837 let local_session_id: Option<SessionId> = if use_server_generated_id {
838 None
839 } else {
840 Some(
841 caller_session_id
842 .clone()
843 .unwrap_or_else(|| SessionId::new(uuid::Uuid::new_v4().to_string())),
844 )
845 };
846 if config.hooks_handler.is_some() && config.hooks.is_none() {
847 config.hooks = Some(true);
848 }
849 if let Some(transforms) = config.system_message_transform.clone() {
850 inject_transform_sections(&mut config, transforms.as_ref());
851 }
852 let mode = self.inner.mode;
853 if mode == crate::ClientMode::Empty && config.available_tools.is_none() {
854 return Err(Error::with_message(
855 ErrorKind::InvalidConfig,
856 "ClientMode::Empty requires available_tools to be set on the session config. \
857 Use ToolSet to specify which tools the session may use (e.g. \
858 ToolSet::new().add_builtin_many(BUILTIN_TOOLS_ISOLATED)).",
859 ));
860 }
861 crate::mode::validate_tool_filter_list(
862 "available_tools",
863 config.available_tools.as_deref(),
864 )?;
865 crate::mode::validate_tool_filter_list("excluded_tools", config.excluded_tools.as_deref())?;
866 config.system_message =
867 crate::mode::system_message_for_mode(mode, config.system_message.take());
868 config.memory = crate::mode::memory_for_mode(mode, config.memory.take());
869 config.enable_experimental_mode =
870 crate::mode::experimental_mode_for_mode(mode, config.enable_experimental_mode);
871 if mode == crate::ClientMode::Empty {
872 if config.enable_session_telemetry.is_none() {
873 config.enable_session_telemetry = Some(false);
874 }
875 if config.skip_embedding_retrieval.is_none() {
876 config.skip_embedding_retrieval = Some(true);
877 }
878 if config.enable_on_demand_instruction_discovery.is_none() {
879 config.enable_on_demand_instruction_discovery = Some(false);
880 }
881 if config.enable_file_hooks.is_none() {
882 config.enable_file_hooks = Some(false);
883 }
884 if config.enable_host_git_operations.is_none() {
885 config.enable_host_git_operations = Some(false);
886 }
887 if config.enable_session_store.is_none() {
888 config.enable_session_store = Some(false);
889 }
890 if config.enable_skills.is_none() {
891 config.enable_skills = Some(false);
892 }
893 }
894 if mode == crate::ClientMode::Empty && config.mcp_oauth_token_storage.is_none() {
895 config.mcp_oauth_token_storage = Some("in-memory".into());
896 }
897 if mode == crate::ClientMode::Empty && config.embedding_cache_storage.is_none() {
898 config.embedding_cache_storage = Some("in-memory".into());
899 }
900 config.custom_agents_local_only =
901 crate::mode::resolve_custom_agents_local_only(mode, config.custom_agents_local_only);
902 let opt_skip_custom_instructions = config.skip_custom_instructions;
903 let opt_custom_agents_local_only = config.custom_agents_local_only;
904 let opt_coauthor_enabled = config.coauthor_enabled;
905 let opt_manage_schedule_enabled = config.manage_schedule_enabled;
906 let (mut wire, mut runtime) = config.into_wire(local_session_id.clone())?;
907 wire.enable_github_telemetry_forwarding =
908 self.inner.on_github_telemetry.is_some().then_some(true);
909
910 let permission_handler = crate::permission::resolve_handler(
911 runtime.permission_handler.take(),
912 runtime.permission_policy.take(),
913 );
914 let handlers = SessionHandlers {
915 permission: permission_handler,
916 managed_settings_enabled: has_managed_settings(
917 wire.enable_managed_settings,
918 wire.managed_settings.as_ref(),
919 ),
920 elicitation: runtime.elicitation_handler.take(),
921 mcp_auth: runtime.mcp_auth_handler.take(),
922 user_input: runtime.user_input_handler.take(),
923 exit_plan_mode: runtime.exit_plan_mode_handler.take(),
924 auto_mode_switch: runtime.auto_mode_switch_handler.take(),
925 tools: Arc::new(std::mem::take(&mut runtime.tool_handlers)),
926 };
927 let hooks = runtime.hooks_handler.take();
928 let transforms = runtime.system_message_transform.take();
929 let tools_count = wire.tools.as_ref().map_or(0, Vec::len);
930 let commands_count = runtime.commands.as_ref().map_or(0, Vec::len);
931 let has_hooks = hooks.is_some();
932 let command_handlers = build_command_handler_map(runtime.commands.as_deref());
933 let canvas_handler = runtime.canvas_handler.take();
934 let session_fs_provider = runtime.session_fs_provider.take();
935 let bearer_token_providers = std::mem::take(&mut runtime.bearer_token_providers);
936 let has_mcp_auth_handler = handlers.mcp_auth.is_some();
937 if self.inner.session_fs_configured && session_fs_provider.is_none() {
938 return Err(ErrorKind::Session(SessionErrorKind::SessionFsProviderRequired).into());
939 }
940 if self.inner.session_fs_sqlite_declared
941 && let Some(ref provider) = session_fs_provider
942 && provider.sqlite().is_none()
943 {
944 return Err(Error::with_message(
945 ErrorKind::InvalidConfig,
946 "SessionFs capabilities declare SQLite support but the provider \
947 does not implement SessionFsSqliteProvider",
948 ));
949 }
950
951 let mut params = serde_json::to_value(&wire)?;
952 let trace_ctx = self.resolve_trace_context().await;
953 inject_trace_context(&mut params, &trace_ctx);
954
955 let setup_start = Instant::now();
956 let capabilities = Arc::new(parking_lot::RwLock::new(SessionCapabilities::default()));
957 let idle_waiter = Arc::new(ParkingLotMutex::new(None));
958 let open_canvases = Arc::new(parking_lot::RwLock::new(Vec::new()));
959 let shutdown = CancellationToken::new();
960 let (event_tx, _) = tokio::sync::broadcast::channel(512);
961
962 let inline_stash: Arc<
968 ParkingLotMutex<Option<(SessionId, crate::router::SessionChannels)>>,
969 > = Arc::new(ParkingLotMutex::new(None));
970
971 let inline_callback: Option<crate::jsonrpc::InlineResponseCallback> = if let Some(ref sid) =
972 local_session_id
973 {
974 let channels = self.register_session(sid);
975 *inline_stash.lock() = Some((sid.clone(), channels));
976 None
977 } else {
978 let client = self.clone();
979 let stash = inline_stash.clone();
980 let expected = caller_session_id.clone();
981 Some(Box::new(move |response| {
982 let result = response.result.as_ref().ok_or_else(|| {
983 Error::with_message(ErrorKind::Json, "session.create response had no result")
984 })?;
985 let parsed: CreateSessionResult =
986 serde_json::from_value(result.clone()).map_err(Error::from)?;
987 if let Some(requested) = expected.as_ref()
988 && parsed.session_id != *requested
989 {
990 return Err(ErrorKind::Session(SessionErrorKind::SessionIdMismatch {
991 requested: requested.clone(),
992 returned: parsed.session_id,
993 })
994 .into());
995 }
996 let channels = client.register_session(&parsed.session_id);
997 *stash.lock() = Some((parsed.session_id, channels));
998 Ok(())
999 }))
1000 };
1001
1002 let rpc_start = Instant::now();
1003 let result = match self
1004 .call_with_inline_callback("session.create", Some(params), inline_callback)
1005 .await
1006 {
1007 Ok(result) => result,
1008 Err(error) => {
1009 if let Some((id, _channels)) = inline_stash.lock().take() {
1010 self.unregister_session(&id);
1011 }
1012 return Err(error);
1013 }
1014 };
1015 tracing::debug!(
1016 elapsed_ms = rpc_start.elapsed().as_millis(),
1017 "Client::create_session session creation request completed successfully"
1018 );
1019 let create_result: CreateSessionResult = match serde_json::from_value(result) {
1020 Ok(result) => result,
1021 Err(error) => {
1022 if let Some((id, _channels)) = inline_stash.lock().take() {
1023 self.unregister_session(&id);
1024 }
1025 return Err(error.into());
1026 }
1027 };
1028
1029 if let Some(ref requested) = local_session_id
1030 && create_result.session_id != *requested
1031 {
1032 if let Some((id, _channels)) = inline_stash.lock().take() {
1033 self.unregister_session(&id);
1034 }
1035 return Err(ErrorKind::Session(SessionErrorKind::SessionIdMismatch {
1036 requested: requested.clone(),
1037 returned: create_result.session_id.clone(),
1038 })
1039 .into());
1040 }
1041
1042 let (session_id, channels) = inline_stash
1043 .lock()
1044 .take()
1045 .expect("session registration must have populated stash on success");
1046 let event_loop = spawn_event_loop(
1047 session_id.clone(),
1048 self.clone(),
1049 handlers,
1050 hooks,
1051 transforms,
1052 command_handlers,
1053 canvas_handler,
1054 session_fs_provider,
1055 bearer_token_providers,
1056 channels,
1057 idle_waiter.clone(),
1058 capabilities.clone(),
1059 open_canvases.clone(),
1060 event_tx.clone(),
1061 shutdown.clone(),
1062 );
1063 tracing::debug!(
1064 elapsed_ms = setup_start.elapsed().as_millis(),
1065 session_id = %session_id,
1066 tools_count,
1067 commands_count,
1068 has_hooks,
1069 "Client::create_session local setup complete"
1070 );
1071 *capabilities.write() = create_result.capabilities.unwrap_or_default();
1072 if has_mcp_auth_handler {
1073 register_mcp_auth_interest(self, &session_id).await?;
1074 }
1075
1076 tracing::debug!(
1077 elapsed_ms = total_start.elapsed().as_millis(),
1078 session_id = %session_id,
1079 "Client::create_session complete"
1080 );
1081 let session = Session {
1082 id: session_id,
1083 cwd: self.cwd().clone(),
1084 workspace_path: create_result.workspace_path,
1085 remote_url: create_result.remote_url,
1086 client: self.clone(),
1087 event_loop: ParkingLotMutex::new(Some(event_loop)),
1088 shutdown,
1089 idle_waiter,
1090 capabilities,
1091 open_canvases,
1092 event_tx,
1093 };
1094 apply_mode_post_create_patch(
1095 &session,
1096 mode,
1097 opt_skip_custom_instructions,
1098 opt_custom_agents_local_only,
1099 opt_coauthor_enabled,
1100 opt_manage_schedule_enabled,
1101 )
1102 .await?;
1103 Ok(session)
1104 }
1105
1106 pub async fn resume_session(&self, mut config: ResumeSessionConfig) -> Result<Session, Error> {
1117 let total_start = Instant::now();
1118 let session_id = config.session_id.clone();
1119 if config.hooks_handler.is_some() && config.hooks.is_none() {
1120 config.hooks = Some(true);
1121 }
1122 if let Some(transforms) = config.system_message_transform.clone() {
1123 inject_transform_sections_resume(&mut config, transforms.as_ref());
1124 }
1125 let mode = self.inner.mode;
1126 if mode == crate::ClientMode::Empty && config.available_tools.is_none() {
1127 return Err(Error::with_message(
1128 ErrorKind::InvalidConfig,
1129 "ClientMode::Empty requires available_tools to be set on the session config. \
1130 Use ToolSet to specify which tools the session may use (e.g. \
1131 ToolSet::new().add_builtin_many(BUILTIN_TOOLS_ISOLATED)).",
1132 ));
1133 }
1134 crate::mode::validate_tool_filter_list(
1135 "available_tools",
1136 config.available_tools.as_deref(),
1137 )?;
1138 crate::mode::validate_tool_filter_list("excluded_tools", config.excluded_tools.as_deref())?;
1139 config.system_message =
1140 crate::mode::system_message_for_mode(mode, config.system_message.take());
1141 config.memory = crate::mode::memory_for_mode(mode, config.memory.take());
1142 config.enable_experimental_mode =
1143 crate::mode::experimental_mode_for_mode(mode, config.enable_experimental_mode);
1144 if mode == crate::ClientMode::Empty {
1145 if config.enable_session_telemetry.is_none() {
1146 config.enable_session_telemetry = Some(false);
1147 }
1148 if config.skip_embedding_retrieval.is_none() {
1149 config.skip_embedding_retrieval = Some(true);
1150 }
1151 if config.enable_on_demand_instruction_discovery.is_none() {
1152 config.enable_on_demand_instruction_discovery = Some(false);
1153 }
1154 if config.enable_file_hooks.is_none() {
1155 config.enable_file_hooks = Some(false);
1156 }
1157 if config.enable_host_git_operations.is_none() {
1158 config.enable_host_git_operations = Some(false);
1159 }
1160 if config.enable_session_store.is_none() {
1161 config.enable_session_store = Some(false);
1162 }
1163 if config.enable_skills.is_none() {
1164 config.enable_skills = Some(false);
1165 }
1166 }
1167 if mode == crate::ClientMode::Empty && config.mcp_oauth_token_storage.is_none() {
1168 config.mcp_oauth_token_storage = Some("in-memory".into());
1169 }
1170 if mode == crate::ClientMode::Empty && config.embedding_cache_storage.is_none() {
1171 config.embedding_cache_storage = Some("in-memory".into());
1172 }
1173 config.custom_agents_local_only =
1174 crate::mode::resolve_custom_agents_local_only(mode, config.custom_agents_local_only);
1175 let opt_skip_custom_instructions = config.skip_custom_instructions;
1176 let opt_custom_agents_local_only = config.custom_agents_local_only;
1177 let opt_coauthor_enabled = config.coauthor_enabled;
1178 let opt_manage_schedule_enabled = config.manage_schedule_enabled;
1179 let (mut wire, mut runtime) = config.into_wire()?;
1180 wire.enable_github_telemetry_forwarding =
1181 self.inner.on_github_telemetry.is_some().then_some(true);
1182
1183 let permission_handler = crate::permission::resolve_handler(
1184 runtime.permission_handler.take(),
1185 runtime.permission_policy.take(),
1186 );
1187 let handlers = SessionHandlers {
1188 permission: permission_handler,
1189 managed_settings_enabled: has_managed_settings(
1190 wire.enable_managed_settings,
1191 wire.managed_settings.as_ref(),
1192 ),
1193 elicitation: runtime.elicitation_handler.take(),
1194 mcp_auth: runtime.mcp_auth_handler.take(),
1195 user_input: runtime.user_input_handler.take(),
1196 exit_plan_mode: runtime.exit_plan_mode_handler.take(),
1197 auto_mode_switch: runtime.auto_mode_switch_handler.take(),
1198 tools: Arc::new(std::mem::take(&mut runtime.tool_handlers)),
1199 };
1200 let hooks = runtime.hooks_handler.take();
1201 let transforms = runtime.system_message_transform.take();
1202 let tools_count = wire.tools.as_ref().map_or(0, Vec::len);
1203 let commands_count = runtime.commands.as_ref().map_or(0, Vec::len);
1204 let has_hooks = hooks.is_some();
1205 let command_handlers = build_command_handler_map(runtime.commands.as_deref());
1206 let canvas_handler = runtime.canvas_handler.take();
1207 let session_fs_provider = runtime.session_fs_provider.take();
1208 let bearer_token_providers = std::mem::take(&mut runtime.bearer_token_providers);
1209 let has_mcp_auth_handler = handlers.mcp_auth.is_some();
1210 if self.inner.session_fs_configured && session_fs_provider.is_none() {
1211 return Err(ErrorKind::Session(SessionErrorKind::SessionFsProviderRequired).into());
1212 }
1213 if self.inner.session_fs_sqlite_declared
1214 && let Some(ref provider) = session_fs_provider
1215 && provider.sqlite().is_none()
1216 {
1217 return Err(Error::with_message(
1218 ErrorKind::InvalidConfig,
1219 "SessionFs capabilities declare SQLite support but the provider \
1220 does not implement SessionFsSqliteProvider",
1221 ));
1222 }
1223
1224 let mut params = serde_json::to_value(&wire)?;
1225 let trace_ctx = self.resolve_trace_context().await;
1226 inject_trace_context(&mut params, &trace_ctx);
1227
1228 let capabilities = Arc::new(parking_lot::RwLock::new(SessionCapabilities::default()));
1229 let setup_start = Instant::now();
1230 let channels = self.register_session(&session_id);
1231 let idle_waiter = Arc::new(ParkingLotMutex::new(None));
1232 let open_canvases = Arc::new(parking_lot::RwLock::new(Vec::new()));
1233 let shutdown = CancellationToken::new();
1234 let (event_tx, _) = tokio::sync::broadcast::channel(512);
1235 let event_loop = spawn_event_loop(
1236 session_id.clone(),
1237 self.clone(),
1238 handlers,
1239 hooks,
1240 transforms,
1241 command_handlers,
1242 canvas_handler,
1243 session_fs_provider,
1244 bearer_token_providers,
1245 channels,
1246 idle_waiter.clone(),
1247 capabilities.clone(),
1248 open_canvases.clone(),
1249 event_tx.clone(),
1250 shutdown.clone(),
1251 );
1252 let mut registration =
1253 PendingSessionRegistration::new(self.clone(), session_id.clone(), shutdown.clone());
1254 tracing::debug!(
1255 elapsed_ms = setup_start.elapsed().as_millis(),
1256 session_id = %session_id,
1257 tools_count,
1258 commands_count,
1259 has_hooks,
1260 "Client::resume_session local setup complete"
1261 );
1262
1263 let rpc_start = Instant::now();
1264 let result = match self.call("session.resume", Some(params)).await {
1265 Ok(result) => result,
1266 Err(error) => {
1267 registration.cleanup(event_loop).await;
1268 return Err(error);
1269 }
1270 };
1271 tracing::debug!(
1272 elapsed_ms = rpc_start.elapsed().as_millis(),
1273 session_id = %session_id,
1274 "Client::resume_session session resume request completed successfully"
1275 );
1276
1277 let resume_result: ResumeSessionResult = match serde_json::from_value(result) {
1278 Ok(result) => result,
1279 Err(error) => {
1280 registration.cleanup(event_loop).await;
1281 return Err(error.into());
1282 }
1283 };
1284 let cli_session_id = resume_result
1285 .session_id
1286 .clone()
1287 .unwrap_or_else(|| session_id.clone());
1288 if cli_session_id != session_id {
1289 registration.cleanup(event_loop).await;
1290 return Err(ErrorKind::Session(SessionErrorKind::SessionIdMismatch {
1291 requested: session_id,
1292 returned: cli_session_id,
1293 })
1294 .into());
1295 }
1296 if let Some(mcp_servers) = wire.mcp_servers.as_ref()
1297 && let Err(error) = self
1298 .call(
1299 "session.mcp.reloadWithConfig",
1300 Some(serde_json::json!({
1301 "sessionId": session_id,
1302 "config": { "mcpServers": mcp_servers },
1303 })),
1304 )
1305 .await
1306 {
1307 registration.cleanup(event_loop).await;
1308 return Err(error);
1309 }
1310 if has_mcp_auth_handler {
1311 register_mcp_auth_interest(self, &session_id).await?;
1312 }
1313
1314 let skills_reload_start = Instant::now();
1316 if let Err(e) = self
1317 .call(
1318 "session.skills.reload",
1319 Some(serde_json::json!({ "sessionId": session_id })),
1320 )
1321 .await
1322 {
1323 warn!(
1324 elapsed_ms = skills_reload_start.elapsed().as_millis(),
1325 session_id = %session_id,
1326 error = %e,
1327 "Client::resume_session skills reload request failed"
1328 );
1329 } else {
1330 tracing::debug!(
1331 elapsed_ms = skills_reload_start.elapsed().as_millis(),
1332 session_id = %session_id,
1333 "Client::resume_session skills reload request completed successfully"
1334 );
1335 }
1336
1337 *capabilities.write() = resume_result.capabilities.unwrap_or_default();
1338 {
1343 let mut snapshots = open_canvases.write();
1344 for snapshot in resume_result.open_canvases.unwrap_or_default() {
1345 upsert_open_canvas_snapshot(&mut snapshots, snapshot);
1346 }
1347 }
1348
1349 tracing::debug!(
1350 elapsed_ms = total_start.elapsed().as_millis(),
1351 session_id = %session_id,
1352 "Client::resume_session complete"
1353 );
1354 registration.disarm();
1355 let session = Session {
1356 id: session_id,
1357 cwd: self.cwd().clone(),
1358 workspace_path: resume_result.workspace_path,
1359 remote_url: resume_result.remote_url,
1360 client: self.clone(),
1361 event_loop: ParkingLotMutex::new(Some(event_loop)),
1362 shutdown,
1363 idle_waiter,
1364 capabilities,
1365 open_canvases,
1366 event_tx,
1367 };
1368 apply_mode_post_create_patch(
1369 &session,
1370 mode,
1371 opt_skip_custom_instructions,
1372 opt_custom_agents_local_only,
1373 opt_coauthor_enabled,
1374 opt_manage_schedule_enabled,
1375 )
1376 .await?;
1377 Ok(session)
1378 }
1379}
1380
1381type CommandHandlerMap = HashMap<String, Arc<dyn CommandHandler>>;
1382
1383async fn apply_mode_post_create_patch(
1384 session: &Session,
1385 mode: crate::ClientMode,
1386 opt_skip_custom_instructions: Option<bool>,
1387 opt_custom_agents_local_only: Option<bool>,
1388 opt_coauthor_enabled: Option<bool>,
1389 opt_manage_schedule_enabled: Option<bool>,
1390) -> Result<(), Error> {
1391 use crate::generated::api_types::SessionUpdateOptionsParams;
1392 let mut patch = SessionUpdateOptionsParams::default();
1393 let should_send = if mode == crate::ClientMode::Empty {
1394 patch.skip_custom_instructions = Some(opt_skip_custom_instructions.unwrap_or(true));
1395 patch.custom_agents_local_only = Some(opt_custom_agents_local_only.unwrap_or(true));
1396 patch.coauthor_enabled = Some(opt_coauthor_enabled.unwrap_or(false));
1397 patch.manage_schedule_enabled = Some(opt_manage_schedule_enabled.unwrap_or(false));
1398 patch.installed_plugins = Some(Vec::new());
1399 true
1400 } else {
1401 let mut any = false;
1402 if let Some(v) = opt_skip_custom_instructions {
1403 patch.skip_custom_instructions = Some(v);
1404 any = true;
1405 }
1406 if let Some(v) = opt_custom_agents_local_only {
1407 patch.custom_agents_local_only = Some(v);
1408 any = true;
1409 }
1410 if let Some(v) = opt_coauthor_enabled {
1411 patch.coauthor_enabled = Some(v);
1412 any = true;
1413 }
1414 if let Some(v) = opt_manage_schedule_enabled {
1415 patch.manage_schedule_enabled = Some(v);
1416 any = true;
1417 }
1418 any
1419 };
1420 if !should_send {
1421 return Ok(());
1422 }
1423 if let Err(error) = session.rpc().options().update(patch).await {
1424 let _ = session.disconnect().await;
1425 return Err(error);
1426 }
1427 Ok(())
1428}
1429
1430fn build_command_handler_map(commands: Option<&[CommandDefinition]>) -> Arc<CommandHandlerMap> {
1431 let map = match commands {
1432 Some(commands) => commands
1433 .iter()
1434 .filter(|cmd| !cmd.name.is_empty())
1435 .map(|cmd| (cmd.name.clone(), cmd.handler.clone()))
1436 .collect(),
1437 None => HashMap::new(),
1438 };
1439 Arc::new(map)
1440}
1441
1442fn upsert_open_canvas_snapshot(
1443 snapshots: &mut Vec<OpenCanvasInstance>,
1444 snapshot: OpenCanvasInstance,
1445) {
1446 if let Some(existing) = snapshots
1447 .iter_mut()
1448 .find(|open| open.instance_id == snapshot.instance_id)
1449 {
1450 *existing = snapshot;
1451 } else {
1452 snapshots.push(snapshot);
1453 }
1454}
1455
1456fn remove_open_canvas_snapshot(snapshots: &mut Vec<OpenCanvasInstance>, instance_id: &str) {
1457 snapshots.retain(|open| open.instance_id != instance_id);
1458}
1459
1460#[allow(clippy::too_many_arguments)]
1461fn spawn_event_loop(
1462 session_id: SessionId,
1463 client: Client,
1464 handlers: SessionHandlers,
1465 hooks: Option<Arc<dyn SessionHooks>>,
1466 transforms: Option<Arc<dyn SystemMessageTransform>>,
1467 command_handlers: Arc<CommandHandlerMap>,
1468 canvas_handler: Option<Arc<dyn CanvasHandler>>,
1469 session_fs_provider: Option<Arc<dyn SessionFsProvider>>,
1470 bearer_token_providers: HashMap<String, Arc<dyn BearerTokenProvider>>,
1471 channels: crate::router::SessionChannels,
1472 idle_waiter: Arc<ParkingLotMutex<Option<IdleWaiter>>>,
1473 capabilities: Arc<parking_lot::RwLock<SessionCapabilities>>,
1474 open_canvases: Arc<parking_lot::RwLock<Vec<OpenCanvasInstance>>>,
1475 event_tx: tokio::sync::broadcast::Sender<SessionEvent>,
1476 shutdown: CancellationToken,
1477) -> JoinHandle<()> {
1478 let crate::router::SessionChannels {
1479 mut notifications,
1480 mut requests,
1481 } = channels;
1482
1483 let span = tracing::error_span!("session_event_loop", session_id = %session_id);
1484 tokio::spawn(
1485 async move {
1486 loop {
1487 tokio::select! {
1511 _ = shutdown.cancelled() => break,
1512 Some(notification) = notifications.recv() => {
1513 handle_notification(
1514 &session_id, &client, &handlers, &command_handlers, notification, &idle_waiter, &capabilities, &open_canvases, &event_tx,
1515 ).await;
1516 }
1517 Some(request) = requests.recv() => {
1518 let span = tracing::error_span!("session_request_handler", session_id = %session_id);
1522 let session_id = session_id.clone();
1523 let client = client.clone();
1524 let handlers = handlers.clone();
1525 let hooks = hooks.clone();
1526 let transforms = transforms.clone();
1527 let canvas_handler = canvas_handler.clone();
1528 let session_fs_provider = session_fs_provider.clone();
1529 let bearer_token_providers = bearer_token_providers.clone();
1530 tokio::spawn(
1531 async move {
1532 let ctx = RequestDispatchContext {
1533 client: &client,
1534 handlers: &handlers,
1535 hooks: hooks.as_deref(),
1536 transforms: transforms.as_deref(),
1537 canvas_handler: canvas_handler.as_ref(),
1538 session_fs_provider: session_fs_provider.as_ref(),
1539 bearer_token_providers: &bearer_token_providers,
1540 };
1541 handle_request(&session_id, ctx, request).await;
1542 }
1543 .instrument(span),
1544 );
1545 }
1546 else => break,
1547 }
1548 }
1549 if let Some(waiter) = idle_waiter.lock().take() {
1552 let _ = waiter
1553 .tx
1554 .send(Err(ErrorKind::Session(SessionErrorKind::EventLoopClosed).into()));
1555 }
1556 }
1557 .instrument(span),
1558 )
1559}
1560
1561fn extract_request_id(data: &Value) -> Option<RequestId> {
1562 data.get("requestId")
1563 .and_then(|v| v.as_str())
1564 .filter(|s| !s.is_empty())
1565 .map(RequestId::new)
1566}
1567
1568fn permission_request_data(
1569 event_data: &Value,
1570 managed_settings_enabled: bool,
1571) -> PermissionRequestData {
1572 let request_data = event_data
1573 .get("permissionRequest")
1574 .cloned()
1575 .unwrap_or_else(|| event_data.clone());
1576 let managed_approval_required = match request_data.get("managedApprovalRequired") {
1577 None => None,
1578 Some(Value::Bool(value)) => Some(*value),
1579 Some(_) => Some(true),
1580 };
1581 match serde_json::from_value::<PermissionRequestData>(request_data) {
1582 Ok(mut data) => {
1583 data.extra = event_data.clone();
1584 data.managed_settings_enabled = managed_settings_enabled;
1585 data
1586 }
1587 Err(_) => PermissionRequestData {
1588 kind: None,
1589 tool_call_id: None,
1590 managed_approval_required,
1591 managed_settings_enabled,
1592 extra: event_data.clone(),
1593 },
1594 }
1595}
1596
1597fn permission_response_params(
1605 session_id: &SessionId,
1606 request_id: &RequestId,
1607 result: &PermissionResult,
1608) -> Option<Value> {
1609 let (decision, decision_context) = match result {
1610 PermissionResult::Decision { decision, context } => (decision, context.clone()),
1611 PermissionResult::NoResult => return None,
1612 };
1613 let mut params = serde_json::to_value(PermissionDecisionRequest {
1614 decision_context,
1615 request_id: request_id.clone(),
1616 result: decision.clone(),
1617 })
1618 .expect("serializing permission response should succeed");
1619 params["sessionId"] =
1620 serde_json::to_value(session_id).expect("serializing session ID should succeed");
1621 Some(params)
1622}
1623
1624async fn register_mcp_auth_interest(client: &Client, session_id: &SessionId) -> Result<(), Error> {
1625 let mut params = serde_json::to_value(RegisterEventInterestParams {
1626 event_type: "mcp.oauth_required".to_string(),
1627 })?;
1628 params["sessionId"] = Value::String(session_id.to_string());
1629 client
1630 .call(rpc_methods::SESSION_EVENTLOG_REGISTERINTEREST, Some(params))
1631 .await?;
1632 Ok(())
1633}
1634
1635fn tool_failure_result(message: impl Into<String>) -> ToolResult {
1636 let message = message.into();
1637 ToolResult::Expanded(ToolResultExpanded {
1638 text_result_for_llm: message.clone(),
1639 result_type: "failure".to_string(),
1640 binary_results_for_llm: None,
1641 session_log: None,
1642 error: Some(message),
1643 tool_telemetry: None,
1644 tool_references: None,
1645 })
1646}
1647
1648#[allow(clippy::too_many_arguments)]
1650async fn handle_notification(
1651 session_id: &SessionId,
1652 client: &Client,
1653 handlers: &SessionHandlers,
1654 command_handlers: &Arc<CommandHandlerMap>,
1655 notification: SessionEventNotification,
1656 idle_waiter: &Arc<ParkingLotMutex<Option<IdleWaiter>>>,
1657 capabilities: &Arc<parking_lot::RwLock<SessionCapabilities>>,
1658 open_canvases: &Arc<parking_lot::RwLock<Vec<OpenCanvasInstance>>>,
1659 event_tx: &tokio::sync::broadcast::Sender<SessionEvent>,
1660) {
1661 let dispatch_start = Instant::now();
1662 let event = notification.event.clone();
1663 let event_type = event.parsed_type();
1664 if event_type == SessionEventType::PermissionRequested {
1665 tracing::debug!(
1666 session_id = %session_id,
1667 event_type = %event.event_type,
1668 "Session::handle_notification permission request received"
1669 );
1670 }
1671
1672 match event_type {
1675 SessionEventType::AssistantMessage
1676 | SessionEventType::SessionIdle
1677 | SessionEventType::SessionError => {
1678 let mut guard = idle_waiter.lock();
1679 if let Some(waiter) = guard.as_mut() {
1680 match event_type {
1681 SessionEventType::AssistantMessage => {
1682 if !waiter.first_assistant_message_seen {
1683 waiter.first_assistant_message_seen = true;
1684 tracing::debug!(
1685 elapsed_ms = waiter.started_at.elapsed().as_millis(),
1686 session_id = %session_id,
1687 "Session::send_and_wait first assistant message"
1688 );
1689 }
1690 waiter.last_assistant_message = Some(event.clone());
1691 }
1692 SessionEventType::SessionIdle | SessionEventType::SessionError => {
1693 if let Some(waiter) = guard.take() {
1694 if event_type == SessionEventType::SessionIdle {
1695 tracing::debug!(
1696 elapsed_ms = waiter.started_at.elapsed().as_millis(),
1697 session_id = %session_id,
1698 "Session::send_and_wait idle received"
1699 );
1700 let _ = waiter.tx.send(Ok(waiter.last_assistant_message));
1701 } else {
1702 let error_msg = event
1703 .typed_data::<SessionErrorData>()
1704 .map(|d| d.message)
1705 .or_else(|| {
1706 event
1707 .data
1708 .get("message")
1709 .and_then(|v| v.as_str())
1710 .map(|s| s.to_string())
1711 })
1712 .unwrap_or_else(|| "session error".to_string());
1713 let _ = waiter.tx.send(Err(Error::with_message(
1714 ErrorKind::Session(SessionErrorKind::AgentError),
1715 error_msg,
1716 )));
1717 }
1718 }
1719 }
1720 _ => {}
1721 }
1722 }
1723 }
1724 _ => {}
1725 }
1726
1727 if event_type == SessionEventType::CapabilitiesChanged {
1731 match serde_json::from_value::<SessionCapabilities>(notification.event.data.clone()) {
1732 Ok(changed) => *capabilities.write() = changed,
1733 Err(e) => warn!(error = %e, "failed to deserialize capabilities.changed payload"),
1734 }
1735 }
1736 if event_type == SessionEventType::SessionCanvasOpened {
1737 match serde_json::from_value::<OpenCanvasInstance>(notification.event.data.clone()) {
1738 Ok(open_canvas) => {
1739 upsert_open_canvas_snapshot(&mut open_canvases.write(), open_canvas);
1740 }
1741 Err(e) => warn!(error = %e, "failed to deserialize session.canvas.opened payload"),
1742 }
1743 }
1744 if event_type == SessionEventType::SessionCanvasClosed {
1745 match serde_json::from_value::<SessionCanvasClosedData>(notification.event.data.clone()) {
1746 Ok(closed) => {
1747 if closed.instance_id.is_empty() {
1748 warn!("failed to deserialize session.canvas.closed payload");
1749 } else {
1750 remove_open_canvas_snapshot(&mut open_canvases.write(), &closed.instance_id);
1751 }
1752 }
1753 Err(e) => warn!(error = %e, "failed to deserialize session.canvas.closed payload"),
1754 }
1755 }
1756
1757 let _ = event_tx.send(event.clone());
1761
1762 tracing::debug!(
1763 elapsed_ms = dispatch_start.elapsed().as_millis(),
1764 session_id = %session_id,
1765 event_type = %notification.event.event_type,
1766 "Session::handle_notification dispatch"
1767 );
1768
1769 match event_type {
1772 SessionEventType::PermissionRequested => {
1773 let Some(request_id) = extract_request_id(¬ification.event.data) else {
1774 return;
1775 };
1776 if notification
1780 .event
1781 .data
1782 .get("resolvedByHook")
1783 .and_then(|v| v.as_bool())
1784 .unwrap_or(false)
1785 {
1786 return;
1787 }
1788 let Some(permission_handler) = handlers.permission.clone() else {
1792 return;
1793 };
1794 let client = client.clone();
1795 let sid = session_id.clone();
1796 let data = permission_request_data(
1797 ¬ification.event.data,
1798 handlers.managed_settings_enabled,
1799 );
1800 let span = tracing::error_span!(
1801 "permission_request_handler",
1802 session_id = %sid,
1803 request_id = %request_id
1804 );
1805 tokio::spawn(
1806 async move {
1807 let handler_start = Instant::now();
1808 let result = permission_handler
1809 .handle(sid.clone(), request_id.clone(), data)
1810 .await;
1811 tracing::debug!(
1812 elapsed_ms = handler_start.elapsed().as_millis(),
1813 session_id = %sid,
1814 request_id = %request_id,
1815 "PermissionHandler::handle dispatch"
1816 );
1817 let Some(params) = permission_response_params(&sid, &request_id, &result)
1818 else {
1819 return;
1823 };
1824 let rpc_start = Instant::now();
1825 let _ = client
1826 .call(
1827 rpc_methods::SESSION_PERMISSIONS_HANDLEPENDINGPERMISSIONREQUEST,
1828 Some(params),
1829 )
1830 .await;
1831 tracing::debug!(
1832 elapsed_ms = rpc_start.elapsed().as_millis(),
1833 session_id = %sid,
1834 request_id = %request_id,
1835 "Session::handle_notification response sent successfully"
1836 );
1837 }
1838 .instrument(span),
1839 );
1840 }
1841 SessionEventType::ExternalToolRequested => {
1842 let Some(request_id) = extract_request_id(¬ification.event.data) else {
1843 return;
1844 };
1845 let data: ExternalToolRequestedData =
1846 match serde_json::from_value(notification.event.data.clone()) {
1847 Ok(d) => d,
1848 Err(e) => {
1849 warn!(error = %e, "failed to deserialize external_tool.requested");
1850 let client = client.clone();
1851 let sid = session_id.clone();
1852 let span = tracing::error_span!(
1853 "external_tool_deserialize_error",
1854 session_id = %sid,
1855 request_id = %request_id
1856 );
1857 tokio::spawn(
1858 async move {
1859 let rpc_start = Instant::now();
1860 let _ = client
1861 .call(
1862 "session.tools.handlePendingToolCall",
1863 Some(serde_json::json!({
1864 "sessionId": sid,
1865 "requestId": request_id,
1866 "error": format!("Failed to deserialize tool request: {e}"),
1867 })),
1868 )
1869 .await;
1870 tracing::debug!(
1871 elapsed_ms = rpc_start.elapsed().as_millis(),
1872 session_id = %sid,
1873 request_id = %request_id,
1874 "Session::handle_notification response sent successfully"
1875 );
1876 }
1877 .instrument(span),
1878 );
1879 return;
1880 }
1881 };
1882 let tool_handler = if data.tool_name.is_empty() {
1886 None
1887 } else {
1888 handlers.tools.get(&data.tool_name).cloned()
1889 };
1890 let Some(tool_handler) = tool_handler else {
1891 return;
1892 };
1893 let client = client.clone();
1894 let sid = session_id.clone();
1895 let span = tracing::error_span!(
1896 "external_tool_handler",
1897 session_id = %sid,
1898 request_id = %request_id
1899 );
1900 tokio::spawn(
1901 async move {
1902 if data.tool_call_id.is_empty() {
1907 let error_msg = "Missing toolCallId";
1908 let rpc_start = Instant::now();
1909 let _ = client
1910 .call(
1911 "session.tools.handlePendingToolCall",
1912 Some(serde_json::json!({
1913 "sessionId": sid,
1914 "requestId": request_id,
1915 "error": error_msg,
1916 })),
1917 )
1918 .await;
1919 tracing::debug!(
1920 elapsed_ms = rpc_start.elapsed().as_millis(),
1921 session_id = %sid,
1922 request_id = %request_id,
1923 "Session::handle_notification response sent successfully"
1924 );
1925 return;
1926 }
1927 let tool_call_id = data.tool_call_id.clone();
1928 let tool_name = data.tool_name.clone();
1929 let available_tools = if tool_name == TOOL_SEARCH_TOOL_NAME {
1936 match client
1937 .call(
1938 rpc_methods::SESSION_TOOLS_GETCURRENTMETADATA,
1939 Some(serde_json::json!({ "sessionId": sid })),
1940 )
1941 .await
1942 {
1943 Ok(value) => {
1944 serde_json::from_value::<ToolsGetCurrentMetadataResult>(value)
1945 .ok()
1946 .and_then(|result| result.tools)
1947 }
1948 Err(_) => None,
1949 }
1950 } else {
1951 None
1952 };
1953 let invocation = ToolInvocation {
1954 session_id: sid.clone(),
1955 tool_call_id: data.tool_call_id,
1956 tool_name: data.tool_name,
1957 arguments: data
1958 .arguments
1959 .unwrap_or(Value::Object(serde_json::Map::new())),
1960 available_tools,
1961 traceparent: data.traceparent,
1962 tracestate: data.tracestate,
1963 };
1964 let handler_start = Instant::now();
1965 let tool_result = match tool_handler.call(invocation).await {
1966 Ok(r) => r,
1967 Err(e) => tool_failure_result(e.to_string()),
1968 };
1969 tracing::debug!(
1970 elapsed_ms = handler_start.elapsed().as_millis(),
1971 session_id = %sid,
1972 request_id = %request_id,
1973 tool_call_id = %tool_call_id,
1974 tool_name = %tool_name,
1975 "ToolHandler::call dispatch"
1976 );
1977 let result_value = serde_json::to_value(tool_result).unwrap_or(Value::Null);
1978 let rpc_start = Instant::now();
1979 let _ = client
1980 .call(
1981 "session.tools.handlePendingToolCall",
1982 Some(serde_json::json!({
1983 "sessionId": sid,
1984 "requestId": request_id,
1985 "result": result_value,
1986 })),
1987 )
1988 .await;
1989 tracing::debug!(
1990 elapsed_ms = rpc_start.elapsed().as_millis(),
1991 session_id = %sid,
1992 request_id = %request_id,
1993 tool_call_id = %tool_call_id,
1994 tool_name = %tool_name,
1995 "Session::handle_notification response sent successfully"
1996 );
1997 }
1998 .instrument(span),
1999 );
2000 }
2001 SessionEventType::UserInputRequested => {
2002 }
2009 SessionEventType::ElicitationRequested => {
2010 let Some(request_id) = extract_request_id(¬ification.event.data) else {
2011 return;
2012 };
2013 let Some(elicitation_handler) = handlers.elicitation.clone() else {
2017 return;
2018 };
2019 let elicitation_data: ElicitationRequestedData =
2020 match serde_json::from_value(notification.event.data.clone()) {
2021 Ok(d) => d,
2022 Err(e) => {
2023 warn!(error = %e, "failed to deserialize elicitation request");
2024 return;
2025 }
2026 };
2027 let request = ElicitationRequest {
2028 message: elicitation_data.message,
2029 requested_schema: elicitation_data
2030 .requested_schema
2031 .map(|s| serde_json::to_value(s).unwrap_or(Value::Null)),
2032 mode: elicitation_data.mode.map(|m| match m {
2033 crate::generated::session_events::ElicitationRequestedMode::Form => {
2034 crate::types::ElicitationMode::Form
2035 }
2036 crate::generated::session_events::ElicitationRequestedMode::Url => {
2037 crate::types::ElicitationMode::Url
2038 }
2039 _ => crate::types::ElicitationMode::Unknown,
2040 }),
2041 elicitation_source: elicitation_data.elicitation_source,
2042 url: elicitation_data.url,
2043 };
2044 let client = client.clone();
2045 let sid = session_id.clone();
2046 let span = tracing::error_span!(
2047 "elicitation_request_handler",
2048 session_id = %sid,
2049 request_id = %request_id
2050 );
2051 tokio::spawn(
2052 async move {
2053 let cancel = ElicitationResult {
2054 action: "cancel".to_string(),
2055 content: None,
2056 };
2057 let handler_task = tokio::spawn({
2059 let sid = sid.clone();
2060 let request_id = request_id.clone();
2061 let span = tracing::error_span!(
2062 "elicitation_callback",
2063 session_id = %sid,
2064 request_id = %request_id
2065 );
2066 async move {
2067 let handler_start = Instant::now();
2068 let response = elicitation_handler
2069 .handle(sid.clone(), request_id.clone(), request)
2070 .await;
2071 tracing::debug!(
2072 elapsed_ms = handler_start.elapsed().as_millis(),
2073 session_id = %sid,
2074 request_id = %request_id,
2075 "ElicitationHandler::handle dispatch"
2076 );
2077 response
2078 }
2079 .instrument(span)
2080 });
2081 let result = match handler_task.await {
2082 Ok(r) => r,
2083 Err(_) => cancel.clone(),
2084 };
2085 let rpc_start = Instant::now();
2086 if let Err(e) = client
2087 .call(
2088 "session.ui.handlePendingElicitation",
2089 Some(serde_json::json!({
2090 "sessionId": sid,
2091 "requestId": request_id,
2092 "result": result,
2093 })),
2094 )
2095 .await
2096 {
2097 warn!(error = %e, "handlePendingElicitation failed, sending cancel");
2099 let _ = client
2100 .call(
2101 "session.ui.handlePendingElicitation",
2102 Some(serde_json::json!({
2103 "sessionId": sid,
2104 "requestId": request_id,
2105 "result": cancel,
2106 })),
2107 )
2108 .await;
2109 } else {
2110 tracing::debug!(
2111 elapsed_ms = rpc_start.elapsed().as_millis(),
2112 session_id = %sid,
2113 request_id = %request_id,
2114 "Session::handle_notification response sent successfully"
2115 );
2116 }
2117 }
2118 .instrument(span),
2119 );
2120 }
2121 SessionEventType::McpOauthRequired => {
2122 let Some(request_id) = extract_request_id(¬ification.event.data) else {
2123 return;
2124 };
2125 let Some(mcp_auth_handler) = handlers.mcp_auth.clone() else {
2126 warn!(
2127 session_id = %session_id,
2128 request_id = %request_id,
2129 "received MCP OAuth request without a registered MCP auth handler"
2130 );
2131 return;
2132 };
2133 let data: McpOauthRequiredData =
2134 match serde_json::from_value(notification.event.data.clone()) {
2135 Ok(d) => d,
2136 Err(e) => {
2137 warn!(error = %e, "failed to deserialize MCP OAuth request");
2138 return;
2139 }
2140 };
2141 let request = McpAuthRequest {
2142 request_id: request_id.clone(),
2143 server_name: data.server_name,
2144 server_url: data.server_url,
2145 reason: data.reason,
2146 www_authenticate_params: data.www_authenticate_params,
2147 resource_metadata: data.resource_metadata,
2148 static_client_config: data.static_client_config,
2149 };
2150 let client = client.clone();
2151 let sid = session_id.clone();
2152 let span = tracing::error_span!(
2153 "mcp_auth_request_handler",
2154 session_id = %sid,
2155 request_id = %request_id
2156 );
2157 tokio::spawn(
2158 async move {
2159 let cancel = McpAuthResult::Cancelled;
2160 let handler_task = tokio::spawn({
2161 let sid = sid.clone();
2162 let request_id = request_id.clone();
2163 let span = tracing::error_span!(
2164 "mcp_auth_callback",
2165 session_id = %sid,
2166 request_id = %request_id
2167 );
2168 async move {
2169 let handler_start = Instant::now();
2170 let response = mcp_auth_handler
2171 .handle(sid.clone(), request_id.clone(), request)
2172 .await;
2173 tracing::debug!(
2174 elapsed_ms = handler_start.elapsed().as_millis(),
2175 session_id = %sid,
2176 request_id = %request_id,
2177 "McpAuthHandler::handle dispatch"
2178 );
2179 response
2180 }
2181 .instrument(span)
2182 });
2183 let result = match handler_task.await {
2184 Ok(result) => result,
2185 Err(_) => cancel,
2186 };
2187 let rpc_start = Instant::now();
2188 let _ = client
2189 .call(
2190 "session.mcp.oauth.handlePendingRequest",
2191 Some(serde_json::json!({
2192 "sessionId": sid,
2193 "requestId": request_id,
2194 "result": result.into_wire(),
2195 })),
2196 )
2197 .await;
2198 tracing::debug!(
2199 elapsed_ms = rpc_start.elapsed().as_millis(),
2200 "Session::handle_notification MCP auth response sent"
2201 );
2202 }
2203 .instrument(span),
2204 );
2205 }
2206 SessionEventType::CommandExecute => {
2207 let data: CommandExecuteData =
2208 match serde_json::from_value(notification.event.data.clone()) {
2209 Ok(d) => d,
2210 Err(e) => {
2211 warn!(error = %e, "failed to deserialize command.execute");
2212 return;
2213 }
2214 };
2215 let client = client.clone();
2216 let command_handlers = command_handlers.clone();
2217 let sid = session_id.clone();
2218 let span = tracing::error_span!("command_handler", session_id = %sid);
2219 tokio::spawn(
2220 async move {
2221 let request_id = data.request_id;
2222 let ack_error = match command_handlers.get(&data.command_name).cloned() {
2223 None => Some(format!("Unknown command: {}", data.command_name)),
2224 Some(handler) => {
2225 let command_name = data.command_name.clone();
2226 let ctx = CommandContext {
2227 session_id: sid.clone(),
2228 command: data.command,
2229 command_name: data.command_name,
2230 args: data.args,
2231 };
2232 let handler_start = Instant::now();
2233 let result = handler.on_command(ctx).await;
2234 tracing::debug!(
2235 elapsed_ms = handler_start.elapsed().as_millis(),
2236 session_id = %sid,
2237 request_id = %request_id,
2238 command_name = %command_name,
2239 "CommandHandler::call dispatch"
2240 );
2241 match result {
2242 Ok(()) => None,
2243 Err(e) => Some(e.to_string()),
2244 }
2245 }
2246 };
2247 let mut params = serde_json::json!({
2248 "sessionId": sid,
2249 "requestId": request_id,
2250 });
2251 if let Some(error_msg) = ack_error {
2252 params["error"] = serde_json::Value::String(error_msg);
2253 }
2254 let rpc_start = Instant::now();
2255 let _ = client
2256 .call("session.commands.handlePendingCommand", Some(params))
2257 .await;
2258 tracing::debug!(
2259 elapsed_ms = rpc_start.elapsed().as_millis(),
2260 session_id = %sid,
2261 request_id = %request_id,
2262 "Session::handle_notification response sent successfully"
2263 );
2264 }
2265 .instrument(span),
2266 );
2267 }
2268 _ => {}
2269 }
2270}
2271
2272struct RequestDispatchContext<'a> {
2273 client: &'a Client,
2274 handlers: &'a SessionHandlers,
2275 hooks: Option<&'a dyn SessionHooks>,
2276 transforms: Option<&'a dyn SystemMessageTransform>,
2277 canvas_handler: Option<&'a Arc<dyn CanvasHandler>>,
2278 session_fs_provider: Option<&'a Arc<dyn SessionFsProvider>>,
2279 bearer_token_providers: &'a HashMap<String, Arc<dyn BearerTokenProvider>>,
2280}
2281
2282async fn handle_request(
2284 session_id: &SessionId,
2285 ctx: RequestDispatchContext<'_>,
2286 request: crate::JsonRpcRequest,
2287) {
2288 let sid = session_id.clone();
2289 let client = ctx.client;
2290 let handlers = ctx.handlers;
2291 let hooks = ctx.hooks;
2292 let transforms = ctx.transforms;
2293 let canvas_handler = ctx.canvas_handler;
2294 let session_fs_provider = ctx.session_fs_provider;
2295 let bearer_token_providers = ctx.bearer_token_providers;
2296
2297 if request.method.starts_with("sessionFs.") {
2298 crate::session_fs_dispatch::dispatch(client, session_fs_provider, request).await;
2299 return;
2300 }
2301
2302 if request.method.starts_with("canvas.") {
2303 crate::canvas_dispatch::dispatch(client, canvas_handler, request).await;
2304 return;
2305 }
2306
2307 if request.method == crate::generated::api_types::rpc_methods::PROVIDERTOKEN_GETTOKEN {
2308 crate::provider_token_dispatch::dispatch(client, bearer_token_providers, request).await;
2309 return;
2310 }
2311
2312 match request.method.as_str() {
2313 "hooks.invoke" => {
2314 let params = request.params.as_ref();
2315 let hook_type = params
2316 .and_then(|p| p.get("hookType"))
2317 .and_then(|v| v.as_str())
2318 .unwrap_or("");
2319 let input = params
2320 .and_then(|p| p.get("input"))
2321 .cloned()
2322 .unwrap_or(Value::Object(Default::default()));
2323
2324 let rpc_result = if let Some(hooks) = hooks {
2325 match crate::hooks::dispatch_hook(hooks, &sid, hook_type, input).await {
2326 Ok(output) => output,
2327 Err(e) => {
2328 warn!(error = %e, hook_type = hook_type, "hook dispatch failed");
2329 serde_json::json!({ "output": {} })
2330 }
2331 }
2332 } else {
2333 serde_json::json!({ "output": {} })
2334 };
2335
2336 let rpc_response = JsonRpcResponse {
2337 jsonrpc: "2.0".to_string(),
2338 id: request.id,
2339 result: Some(rpc_result),
2340 error: None,
2341 };
2342 let _ = client.send_response(&rpc_response).await;
2343 }
2344
2345 "userInput.request" => {
2346 let params = request.params.as_ref();
2347 let Some(question) = params
2348 .and_then(|p| p.get("question"))
2349 .and_then(|v| v.as_str())
2350 else {
2351 warn!("userInput.request missing 'question' field");
2352 let rpc_response = JsonRpcResponse {
2353 jsonrpc: "2.0".to_string(),
2354 id: request.id,
2355 result: None,
2356 error: Some(crate::JsonRpcError {
2357 code: error_codes::INVALID_PARAMS,
2358 message: "missing required field: question".to_string(),
2359 data: None,
2360 }),
2361 };
2362 let _ = client.send_response(&rpc_response).await;
2363 return;
2364 };
2365 let question = question.to_string();
2366 let choices = params
2367 .and_then(|p| p.get("choices"))
2368 .and_then(|v| v.as_array())
2369 .map(|arr| {
2370 arr.iter()
2371 .filter_map(|v| v.as_str().map(|s| s.to_string()))
2372 .collect()
2373 });
2374 let allow_freeform = params
2375 .and_then(|p| p.get("allowFreeform"))
2376 .and_then(|v| v.as_bool());
2377
2378 let handler_start = Instant::now();
2379 let response = if let Some(user_input_handler) = handlers.user_input.as_ref() {
2380 user_input_handler
2381 .handle(sid.clone(), question, choices, allow_freeform)
2382 .await
2383 } else {
2384 None
2385 };
2386 tracing::debug!(
2387 elapsed_ms = handler_start.elapsed().as_millis(),
2388 session_id = %sid,
2389 "UserInputHandler::handle dispatch"
2390 );
2391
2392 let rpc_result = match response {
2393 Some(UserInputResponse {
2394 answer,
2395 was_freeform,
2396 }) => serde_json::json!({
2397 "answer": answer,
2398 "wasFreeform": was_freeform,
2399 }),
2400 None => serde_json::json!({ "noResponse": true }),
2401 };
2402 let rpc_response = JsonRpcResponse {
2403 jsonrpc: "2.0".to_string(),
2404 id: request.id,
2405 result: Some(rpc_result),
2406 error: None,
2407 };
2408 let _ = client.send_response(&rpc_response).await;
2409 }
2410
2411 "exitPlanMode.request" => {
2412 let params = request
2413 .params
2414 .as_ref()
2415 .cloned()
2416 .unwrap_or(Value::Object(serde_json::Map::new()));
2417 let data: ExitPlanModeData = match serde_json::from_value(params) {
2418 Ok(d) => d,
2419 Err(e) => {
2420 warn!(error = %e, "failed to deserialize exitPlanMode.request params, using defaults");
2421 ExitPlanModeData::default()
2422 }
2423 };
2424
2425 let rpc_result = if let Some(exit_plan_handler) = handlers.exit_plan_mode.as_ref() {
2426 let result = exit_plan_handler.handle(sid, data).await;
2427 serde_json::to_value(result).expect("ExitPlanModeResult serialization cannot fail")
2428 } else {
2429 serde_json::json!({ "approved": true })
2430 };
2431 let rpc_response = JsonRpcResponse {
2432 jsonrpc: "2.0".to_string(),
2433 id: request.id,
2434 result: Some(rpc_result),
2435 error: None,
2436 };
2437 let _ = client.send_response(&rpc_response).await;
2438 }
2439
2440 "autoModeSwitch.request" => {
2441 let error_code = request
2442 .params
2443 .as_ref()
2444 .and_then(|p| p.get("errorCode"))
2445 .and_then(|v| v.as_str())
2446 .map(|s| s.to_string());
2447 let retry_after_seconds = request
2448 .params
2449 .as_ref()
2450 .and_then(|p| p.get("retryAfterSeconds"))
2451 .and_then(|v| v.as_f64());
2452
2453 let answer = if let Some(auto_mode_handler) = handlers.auto_mode_switch.as_ref() {
2454 auto_mode_handler
2455 .handle(sid, error_code, retry_after_seconds)
2456 .await
2457 } else {
2458 AutoModeSwitchResponse::No
2459 };
2460 let rpc_response = JsonRpcResponse {
2461 jsonrpc: "2.0".to_string(),
2462 id: request.id,
2463 result: Some(serde_json::json!({ "response": answer })),
2464 error: None,
2465 };
2466 let _ = client.send_response(&rpc_response).await;
2467 }
2468
2469 "systemMessage.transform" => {
2470 let params = request.params.as_ref();
2471 let sections: HashMap<String, crate::transforms::TransformSection> =
2472 match params.and_then(|p| p.get("sections")) {
2473 Some(v) => match serde_json::from_value(v.clone()) {
2474 Ok(s) => s,
2475 Err(e) => {
2476 let _ = send_error_response(
2477 client,
2478 request.id,
2479 error_codes::INVALID_PARAMS,
2480 &format!("invalid sections: {e}"),
2481 )
2482 .await;
2483 return;
2484 }
2485 },
2486 None => {
2487 let _ = send_error_response(
2488 client,
2489 request.id,
2490 error_codes::INVALID_PARAMS,
2491 "missing sections parameter",
2492 )
2493 .await;
2494 return;
2495 }
2496 };
2497
2498 let rpc_result = if let Some(transforms) = transforms {
2499 let transform_start = Instant::now();
2500 let response =
2501 crate::transforms::dispatch_transform(transforms, &sid, sections).await;
2502 tracing::debug!(
2503 elapsed_ms = transform_start.elapsed().as_millis(),
2504 session_id = %sid,
2505 "SystemMessageTransform::transform_section dispatch"
2506 );
2507 match serde_json::to_value(response) {
2508 Ok(v) => v,
2509 Err(e) => {
2510 warn!(error = %e, "failed to serialize transform response");
2511 serde_json::json!({ "sections": {} })
2512 }
2513 }
2514 } else {
2515 let passthrough: HashMap<String, crate::transforms::TransformSection> = sections;
2517 serde_json::json!({ "sections": passthrough })
2518 };
2519
2520 let rpc_response = JsonRpcResponse {
2521 jsonrpc: "2.0".to_string(),
2522 id: request.id,
2523 result: Some(rpc_result),
2524 error: None,
2525 };
2526 let _ = client.send_response(&rpc_response).await;
2527 }
2528
2529 method => {
2530 warn!(
2531 method = method,
2532 "unhandled request method in session event loop"
2533 );
2534 let _ = send_error_response(
2535 client,
2536 request.id,
2537 error_codes::METHOD_NOT_FOUND,
2538 &format!("unknown method: {method}"),
2539 )
2540 .await;
2541 }
2542 }
2543}
2544
2545async fn send_error_response(
2546 client: &Client,
2547 id: u64,
2548 code: i32,
2549 message: &str,
2550) -> Result<(), Error> {
2551 let response = JsonRpcResponse {
2552 jsonrpc: "2.0".to_string(),
2553 id,
2554 result: None,
2555 error: Some(crate::JsonRpcError {
2556 code,
2557 message: message.to_string(),
2558 data: None,
2559 }),
2560 };
2561 client.send_response(&response).await
2562}
2563
2564fn apply_transform_sections(
2568 sys_msg: &mut SystemMessageConfig,
2569 transforms: &dyn SystemMessageTransform,
2570) {
2571 sys_msg.mode = Some("customize".to_string());
2572 let sections = sys_msg.sections.get_or_insert_with(HashMap::new);
2573 for id in transforms.section_ids() {
2574 sections.entry(id).or_insert_with(|| SectionOverride {
2575 action: Some("transform".to_string()),
2576 content: None,
2577 });
2578 }
2579}
2580
2581fn inject_transform_sections(config: &mut SessionConfig, transforms: &dyn SystemMessageTransform) {
2582 let sys_msg = config.system_message.get_or_insert_with(Default::default);
2583 apply_transform_sections(sys_msg, transforms);
2584}
2585
2586fn inject_transform_sections_resume(
2587 config: &mut ResumeSessionConfig,
2588 transforms: &dyn SystemMessageTransform,
2589) {
2590 let sys_msg = config.system_message.get_or_insert_with(Default::default);
2591 apply_transform_sections(sys_msg, transforms);
2592}
2593
2594#[cfg(test)]
2595mod tests {
2596 use serde_json::json;
2597
2598 use super::{has_managed_settings, permission_request_data, permission_response_params};
2599 use crate::handler::PermissionResult;
2600 use crate::types::{
2601 PermissionDecisionContext, PermissionDecisionOutcome, PermissionDecisionSource,
2602 PermissionDecisionSurface, RequestId, SessionId,
2603 };
2604
2605 #[test]
2606 fn direct_injection_enables_managed_safeguards() {
2607 let settings = crate::types::ManagedSettings::default();
2608 assert!(has_managed_settings(None, Some(&settings)));
2609 assert!(!has_managed_settings(None, None));
2610 }
2611
2612 fn attribution_context() -> PermissionDecisionContext {
2613 PermissionDecisionContext {
2614 outcome: PermissionDecisionOutcome::AutoApproved,
2615 source: PermissionDecisionSource::AssistedApproval,
2616 surface: PermissionDecisionSurface::CopilotApp,
2617 }
2618 }
2619
2620 #[test]
2621 fn response_params_omit_decision_context_without_attribution() {
2622 for (result, expected) in [
2623 (
2624 PermissionResult::approve_once(),
2625 json!({ "kind": "approve-once" }),
2626 ),
2627 (PermissionResult::reject(None), json!({ "kind": "reject" })),
2628 (
2629 PermissionResult::reject(Some("bad".to_string())),
2630 json!({ "kind": "reject", "feedback": "bad" }),
2631 ),
2632 (
2633 PermissionResult::user_not_available(),
2634 json!({ "kind": "user-not-available" }),
2635 ),
2636 ] {
2637 let params = permission_response_params(
2638 &SessionId::from("session-1"),
2639 &RequestId::from("permission-1"),
2640 &result,
2641 )
2642 .unwrap();
2643 assert_eq!(
2644 params,
2645 json!({
2646 "sessionId": "session-1",
2647 "requestId": "permission-1",
2648 "result": expected,
2649 })
2650 );
2651 }
2652 }
2653
2654 #[test]
2655 fn response_params_forward_decision_context_alongside_result() {
2656 let params = permission_response_params(
2657 &SessionId::from("session-1"),
2658 &RequestId::from("permission-1"),
2659 &PermissionResult::approve_once().with_context(attribution_context()),
2660 )
2661 .unwrap();
2662 assert_eq!(
2663 params,
2664 json!({
2665 "sessionId": "session-1",
2666 "requestId": "permission-1",
2667 "result": { "kind": "approve-once" },
2668 "decisionContext": {
2669 "outcome": "auto_approved",
2670 "source": "assisted_approval",
2671 "surface": "copilot_app",
2672 },
2673 })
2674 );
2675 assert!(params["result"].get("decisionContext").is_none());
2677 }
2678
2679 #[test]
2680 fn response_params_suppressed_for_no_result() {
2681 assert!(
2682 permission_response_params(
2683 &SessionId::from("session-1"),
2684 &RequestId::from("permission-1"),
2685 &PermissionResult::NoResult,
2686 )
2687 .is_none()
2688 );
2689 }
2690
2691 #[test]
2692 fn with_context_is_a_no_op_on_no_result() {
2693 let result = PermissionResult::no_result().with_context(attribution_context());
2694 assert!(matches!(result, PermissionResult::NoResult));
2695 }
2696
2697 #[test]
2698 fn with_context_replaces_rather_than_nests() {
2699 let result = PermissionResult::approve_once()
2700 .with_context(attribution_context())
2701 .with_context(PermissionDecisionContext {
2702 outcome: PermissionDecisionOutcome::PromptedUser,
2703 source: PermissionDecisionSource::HumanResponse,
2704 surface: PermissionDecisionSurface::Sdk,
2705 });
2706 let params = permission_response_params(
2707 &SessionId::from("session-1"),
2708 &RequestId::from("permission-1"),
2709 &result,
2710 )
2711 .unwrap();
2712 assert_eq!(
2713 params["decisionContext"],
2714 json!({
2715 "outcome": "prompted_user",
2716 "source": "human_response",
2717 "surface": "sdk",
2718 })
2719 );
2720 }
2721
2722 #[test]
2723 fn permission_request_data_reads_nested_managed_approval_metadata() {
2724 let data = permission_request_data(
2725 &json!({
2726 "requestId": "permission-1",
2727 "permissionRequest": {
2728 "kind": "read",
2729 "managedApprovalRequired": true,
2730 "path": "/workspace/file.txt"
2731 }
2732 }),
2733 false,
2734 );
2735
2736 assert_eq!(data.managed_approval_required, Some(true));
2737 assert_eq!(
2738 data.extra["permissionRequest"]["path"],
2739 "/workspace/file.txt"
2740 );
2741 }
2742
2743 #[test]
2744 fn permission_request_data_preserves_managed_flag_when_other_fields_are_malformed() {
2745 let data = permission_request_data(
2746 &json!({
2747 "requestId": "permission-1",
2748 "permissionRequest": {
2749 "kind": "read",
2750 "managedApprovalRequired": true,
2751 "toolCallId": 42
2752 }
2753 }),
2754 false,
2755 );
2756
2757 assert_eq!(data.managed_approval_required, Some(true));
2758 assert_eq!(data.extra["requestId"], "permission-1");
2759 }
2760
2761 #[test]
2762 fn permission_request_data_fails_closed_for_malformed_managed_flag() {
2763 let data = permission_request_data(
2764 &json!({
2765 "requestId": "permission-1",
2766 "permissionRequest": {
2767 "kind": "read",
2768 "managedApprovalRequired": "yes",
2769 "path": "/workspace/file.txt"
2770 }
2771 }),
2772 false,
2773 );
2774
2775 assert_eq!(data.managed_approval_required, Some(true));
2776 }
2777
2778 #[test]
2779 fn permission_request_data_preserves_valid_false_managed_flag() {
2780 let data = permission_request_data(
2781 &json!({
2782 "requestId": "permission-1",
2783 "permissionRequest": {
2784 "kind": "read",
2785 "managedApprovalRequired": false,
2786 "path": "/workspace/file.txt"
2787 }
2788 }),
2789 false,
2790 );
2791
2792 assert_eq!(data.managed_approval_required, Some(false));
2793 }
2794}