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 model_id: model.to_string(),
544 reasoning_effort: opts.reasoning_effort,
545 reasoning_summary: opts.reasoning_summary,
546 verbosity: None,
547 context_tier: opts.context_tier,
548 model_capabilities: opts.model_capabilities,
549 defer_if_model_change_queued: None,
550 };
551 self.rpc().model().switch_to(request).await?;
552 Ok(())
553 }
554
555 pub async fn disconnect(&self) -> Result<(), Error> {
572 self.client
573 .call(
574 "session.destroy",
575 Some(serde_json::json!({ "sessionId": self.id })),
576 )
577 .await?;
578 self.stop_event_loop().await;
579 self.client.unregister_session(&self.id);
580 Ok(())
581 }
582
583 #[deprecated(since = "0.1.0", note = "Use `disconnect()` instead")]
588 pub async fn destroy(&self) -> Result<(), Error> {
589 self.disconnect().await
590 }
591
592 pub async fn log(
596 &self,
597 message: &str,
598 opts: Option<crate::types::LogOptions>,
599 ) -> Result<(), Error> {
600 let opts = opts.unwrap_or_default();
601 let level = match opts.level {
602 Some(level) => Some(serde_json::from_value(serde_json::to_value(level)?)?),
603 None => None,
604 };
605 let request = LogRequest {
606 message: message.to_string(),
607 level,
608 ephemeral: opts.ephemeral,
609 r#type: None,
610 tip: None,
611 url: None,
612 };
613 self.rpc().log(request).await?;
614 Ok(())
615 }
616
617 pub fn ui(&self) -> SessionUi<'_> {
623 SessionUi { session: self }
624 }
625
626 fn assert_elicitation(&self) -> Result<(), Error> {
628 if self
629 .capabilities
630 .read()
631 .ui
632 .as_ref()
633 .and_then(|u| u.elicitation)
634 != Some(true)
635 {
636 return Err(ErrorKind::Session(SessionErrorKind::ElicitationNotSupported).into());
637 }
638 Ok(())
639 }
640}
641
642impl Drop for Session {
643 fn drop(&mut self) {
644 self.shutdown.cancel();
656 self.client.unregister_session(&self.id);
657 }
658}
659
660pub struct SessionUi<'a> {
667 session: &'a Session,
668}
669
670impl<'a> SessionUi<'a> {
671 pub async fn elicitation(
679 &self,
680 message: &str,
681 schema: Value,
682 ) -> Result<ElicitationResult, Error> {
683 self.session.assert_elicitation()?;
684 let result = self
685 .session
686 .client
687 .call(
688 "session.ui.elicitation",
689 Some(serde_json::json!({
690 "sessionId": self.session.id,
691 "message": message,
692 "requestedSchema": schema,
693 })),
694 )
695 .await?;
696 let elicitation: ElicitationResult = serde_json::from_value(result)?;
697 Ok(elicitation)
698 }
699
700 pub async fn confirm(&self, message: &str) -> Result<bool, Error> {
704 self.session.assert_elicitation()?;
705 let schema = serde_json::json!({
706 "type": "object",
707 "properties": {
708 "confirmed": {
709 "type": "boolean",
710 "default": true,
711 }
712 },
713 "required": ["confirmed"]
714 });
715 let result = self.elicitation(message, schema).await?;
716 Ok(result.action == "accept"
717 && result
718 .content
719 .and_then(|c| c.get("confirmed").and_then(|v| v.as_bool()))
720 == Some(true))
721 }
722
723 pub async fn select(&self, message: &str, options: &[&str]) -> Result<Option<String>, Error> {
727 self.session.assert_elicitation()?;
728 let schema = serde_json::json!({
729 "type": "object",
730 "properties": {
731 "selection": {
732 "type": "string",
733 "enum": options,
734 }
735 },
736 "required": ["selection"]
737 });
738 let result = self.elicitation(message, schema).await?;
739 if result.action != "accept" {
740 return Ok(None);
741 }
742 let selection = result.content.and_then(|c| {
743 c.get("selection")
744 .and_then(|v| v.as_str())
745 .map(String::from)
746 });
747 Ok(selection)
748 }
749
750 pub async fn input(
755 &self,
756 message: &str,
757 options: Option<&UiInputOptions<'_>>,
758 ) -> Result<Option<String>, Error> {
759 self.session.assert_elicitation()?;
760 let mut field = serde_json::json!({ "type": "string" });
761 if let Some(opts) = options {
762 if let Some(title) = opts.title {
763 field["title"] = Value::String(title.to_string());
764 }
765 if let Some(desc) = opts.description {
766 field["description"] = Value::String(desc.to_string());
767 }
768 if let Some(min) = opts.min_length {
769 field["minLength"] = Value::Number(min.into());
770 }
771 if let Some(max) = opts.max_length {
772 field["maxLength"] = Value::Number(max.into());
773 }
774 if let Some(fmt) = &opts.format {
775 field["format"] = Value::String(fmt.as_str().to_string());
776 }
777 if let Some(default) = opts.default {
778 field["default"] = Value::String(default.to_string());
779 }
780 }
781 let schema = serde_json::json!({
782 "type": "object",
783 "properties": { "value": field },
784 "required": ["value"]
785 });
786 let result = self.elicitation(message, schema).await?;
787 if result.action != "accept" {
788 return Ok(None);
789 }
790 let value = result
791 .content
792 .and_then(|c| c.get("value").and_then(|v| v.as_str()).map(String::from));
793 Ok(value)
794 }
795}
796
797impl Client {
798 pub async fn create_session(&self, mut config: SessionConfig) -> Result<Session, Error> {
820 let total_start = Instant::now();
821 let caller_session_id = config.session_id.clone();
829 let use_server_generated_id = config.cloud.is_some() && caller_session_id.is_none();
830 let local_session_id: Option<SessionId> = if use_server_generated_id {
831 None
832 } else {
833 Some(
834 caller_session_id
835 .clone()
836 .unwrap_or_else(|| SessionId::new(uuid::Uuid::new_v4().to_string())),
837 )
838 };
839 if config.hooks_handler.is_some() && config.hooks.is_none() {
840 config.hooks = Some(true);
841 }
842 if let Some(transforms) = config.system_message_transform.clone() {
843 inject_transform_sections(&mut config, transforms.as_ref());
844 }
845 let mode = self.inner.mode;
846 if mode == crate::ClientMode::Empty && config.available_tools.is_none() {
847 return Err(Error::with_message(
848 ErrorKind::InvalidConfig,
849 "ClientMode::Empty requires available_tools to be set on the session config. \
850 Use ToolSet to specify which tools the session may use (e.g. \
851 ToolSet::new().add_builtin_many(BUILTIN_TOOLS_ISOLATED)).",
852 ));
853 }
854 crate::mode::validate_tool_filter_list(
855 "available_tools",
856 config.available_tools.as_deref(),
857 )?;
858 crate::mode::validate_tool_filter_list("excluded_tools", config.excluded_tools.as_deref())?;
859 config.system_message =
860 crate::mode::system_message_for_mode(mode, config.system_message.take());
861 config.memory = crate::mode::memory_for_mode(mode, config.memory.take());
862 config.enable_experimental_mode =
863 crate::mode::experimental_mode_for_mode(mode, config.enable_experimental_mode);
864 if mode == crate::ClientMode::Empty {
865 if config.enable_session_telemetry.is_none() {
866 config.enable_session_telemetry = Some(false);
867 }
868 if config.skip_embedding_retrieval.is_none() {
869 config.skip_embedding_retrieval = Some(true);
870 }
871 if config.enable_on_demand_instruction_discovery.is_none() {
872 config.enable_on_demand_instruction_discovery = Some(false);
873 }
874 if config.enable_file_hooks.is_none() {
875 config.enable_file_hooks = Some(false);
876 }
877 if config.enable_host_git_operations.is_none() {
878 config.enable_host_git_operations = Some(false);
879 }
880 if config.enable_session_store.is_none() {
881 config.enable_session_store = Some(false);
882 }
883 if config.enable_skills.is_none() {
884 config.enable_skills = Some(false);
885 }
886 }
887 if mode == crate::ClientMode::Empty && config.mcp_oauth_token_storage.is_none() {
888 config.mcp_oauth_token_storage = Some("in-memory".into());
889 }
890 if mode == crate::ClientMode::Empty && config.embedding_cache_storage.is_none() {
891 config.embedding_cache_storage = Some("in-memory".into());
892 }
893 config.custom_agents_local_only =
894 crate::mode::resolve_custom_agents_local_only(mode, config.custom_agents_local_only);
895 let opt_skip_custom_instructions = config.skip_custom_instructions;
896 let opt_custom_agents_local_only = config.custom_agents_local_only;
897 let opt_coauthor_enabled = config.coauthor_enabled;
898 let opt_manage_schedule_enabled = config.manage_schedule_enabled;
899 let (mut wire, mut runtime) = config.into_wire(local_session_id.clone())?;
900 wire.enable_github_telemetry_forwarding =
901 self.inner.on_github_telemetry.is_some().then_some(true);
902
903 let permission_handler = crate::permission::resolve_handler(
904 runtime.permission_handler.take(),
905 runtime.permission_policy.take(),
906 );
907 let handlers = SessionHandlers {
908 permission: permission_handler,
909 managed_settings_enabled: has_managed_settings(
910 wire.enable_managed_settings,
911 wire.managed_settings.as_ref(),
912 ),
913 elicitation: runtime.elicitation_handler.take(),
914 mcp_auth: runtime.mcp_auth_handler.take(),
915 user_input: runtime.user_input_handler.take(),
916 exit_plan_mode: runtime.exit_plan_mode_handler.take(),
917 auto_mode_switch: runtime.auto_mode_switch_handler.take(),
918 tools: Arc::new(std::mem::take(&mut runtime.tool_handlers)),
919 };
920 let hooks = runtime.hooks_handler.take();
921 let transforms = runtime.system_message_transform.take();
922 let tools_count = wire.tools.as_ref().map_or(0, Vec::len);
923 let commands_count = runtime.commands.as_ref().map_or(0, Vec::len);
924 let has_hooks = hooks.is_some();
925 let command_handlers = build_command_handler_map(runtime.commands.as_deref());
926 let canvas_handler = runtime.canvas_handler.take();
927 let session_fs_provider = runtime.session_fs_provider.take();
928 let bearer_token_providers = std::mem::take(&mut runtime.bearer_token_providers);
929 let has_mcp_auth_handler = handlers.mcp_auth.is_some();
930 if self.inner.session_fs_configured && session_fs_provider.is_none() {
931 return Err(ErrorKind::Session(SessionErrorKind::SessionFsProviderRequired).into());
932 }
933 if self.inner.session_fs_sqlite_declared
934 && let Some(ref provider) = session_fs_provider
935 && provider.sqlite().is_none()
936 {
937 return Err(Error::with_message(
938 ErrorKind::InvalidConfig,
939 "SessionFs capabilities declare SQLite support but the provider \
940 does not implement SessionFsSqliteProvider",
941 ));
942 }
943
944 let mut params = serde_json::to_value(&wire)?;
945 let trace_ctx = self.resolve_trace_context().await;
946 inject_trace_context(&mut params, &trace_ctx);
947
948 let setup_start = Instant::now();
949 let capabilities = Arc::new(parking_lot::RwLock::new(SessionCapabilities::default()));
950 let idle_waiter = Arc::new(ParkingLotMutex::new(None));
951 let open_canvases = Arc::new(parking_lot::RwLock::new(Vec::new()));
952 let shutdown = CancellationToken::new();
953 let (event_tx, _) = tokio::sync::broadcast::channel(512);
954
955 let inline_stash: Arc<
961 ParkingLotMutex<Option<(SessionId, crate::router::SessionChannels)>>,
962 > = Arc::new(ParkingLotMutex::new(None));
963
964 let inline_callback: Option<crate::jsonrpc::InlineResponseCallback> = if let Some(ref sid) =
965 local_session_id
966 {
967 let channels = self.register_session(sid);
968 *inline_stash.lock() = Some((sid.clone(), channels));
969 None
970 } else {
971 let client = self.clone();
972 let stash = inline_stash.clone();
973 let expected = caller_session_id.clone();
974 Some(Box::new(move |response| {
975 let result = response.result.as_ref().ok_or_else(|| {
976 Error::with_message(ErrorKind::Json, "session.create response had no result")
977 })?;
978 let parsed: CreateSessionResult =
979 serde_json::from_value(result.clone()).map_err(Error::from)?;
980 if let Some(requested) = expected.as_ref()
981 && parsed.session_id != *requested
982 {
983 return Err(ErrorKind::Session(SessionErrorKind::SessionIdMismatch {
984 requested: requested.clone(),
985 returned: parsed.session_id,
986 })
987 .into());
988 }
989 let channels = client.register_session(&parsed.session_id);
990 *stash.lock() = Some((parsed.session_id, channels));
991 Ok(())
992 }))
993 };
994
995 let rpc_start = Instant::now();
996 let result = match self
997 .call_with_inline_callback("session.create", Some(params), inline_callback)
998 .await
999 {
1000 Ok(result) => result,
1001 Err(error) => {
1002 if let Some((id, _channels)) = inline_stash.lock().take() {
1003 self.unregister_session(&id);
1004 }
1005 return Err(error);
1006 }
1007 };
1008 tracing::debug!(
1009 elapsed_ms = rpc_start.elapsed().as_millis(),
1010 "Client::create_session session creation request completed successfully"
1011 );
1012 let create_result: CreateSessionResult = match serde_json::from_value(result) {
1013 Ok(result) => result,
1014 Err(error) => {
1015 if let Some((id, _channels)) = inline_stash.lock().take() {
1016 self.unregister_session(&id);
1017 }
1018 return Err(error.into());
1019 }
1020 };
1021
1022 if let Some(ref requested) = local_session_id
1023 && create_result.session_id != *requested
1024 {
1025 if let Some((id, _channels)) = inline_stash.lock().take() {
1026 self.unregister_session(&id);
1027 }
1028 return Err(ErrorKind::Session(SessionErrorKind::SessionIdMismatch {
1029 requested: requested.clone(),
1030 returned: create_result.session_id.clone(),
1031 })
1032 .into());
1033 }
1034
1035 let (session_id, channels) = inline_stash
1036 .lock()
1037 .take()
1038 .expect("session registration must have populated stash on success");
1039 let event_loop = spawn_event_loop(
1040 session_id.clone(),
1041 self.clone(),
1042 handlers,
1043 hooks,
1044 transforms,
1045 command_handlers,
1046 canvas_handler,
1047 session_fs_provider,
1048 bearer_token_providers,
1049 channels,
1050 idle_waiter.clone(),
1051 capabilities.clone(),
1052 open_canvases.clone(),
1053 event_tx.clone(),
1054 shutdown.clone(),
1055 );
1056 tracing::debug!(
1057 elapsed_ms = setup_start.elapsed().as_millis(),
1058 session_id = %session_id,
1059 tools_count,
1060 commands_count,
1061 has_hooks,
1062 "Client::create_session local setup complete"
1063 );
1064 *capabilities.write() = create_result.capabilities.unwrap_or_default();
1065 if has_mcp_auth_handler {
1066 register_mcp_auth_interest(self, &session_id).await?;
1067 }
1068
1069 tracing::debug!(
1070 elapsed_ms = total_start.elapsed().as_millis(),
1071 session_id = %session_id,
1072 "Client::create_session complete"
1073 );
1074 let session = Session {
1075 id: session_id,
1076 cwd: self.cwd().clone(),
1077 workspace_path: create_result.workspace_path,
1078 remote_url: create_result.remote_url,
1079 client: self.clone(),
1080 event_loop: ParkingLotMutex::new(Some(event_loop)),
1081 shutdown,
1082 idle_waiter,
1083 capabilities,
1084 open_canvases,
1085 event_tx,
1086 };
1087 apply_mode_post_create_patch(
1088 &session,
1089 mode,
1090 opt_skip_custom_instructions,
1091 opt_custom_agents_local_only,
1092 opt_coauthor_enabled,
1093 opt_manage_schedule_enabled,
1094 )
1095 .await?;
1096 Ok(session)
1097 }
1098
1099 pub async fn resume_session(&self, mut config: ResumeSessionConfig) -> Result<Session, Error> {
1110 let total_start = Instant::now();
1111 let session_id = config.session_id.clone();
1112 if config.hooks_handler.is_some() && config.hooks.is_none() {
1113 config.hooks = Some(true);
1114 }
1115 if let Some(transforms) = config.system_message_transform.clone() {
1116 inject_transform_sections_resume(&mut config, transforms.as_ref());
1117 }
1118 let mode = self.inner.mode;
1119 if mode == crate::ClientMode::Empty && config.available_tools.is_none() {
1120 return Err(Error::with_message(
1121 ErrorKind::InvalidConfig,
1122 "ClientMode::Empty requires available_tools to be set on the session config. \
1123 Use ToolSet to specify which tools the session may use (e.g. \
1124 ToolSet::new().add_builtin_many(BUILTIN_TOOLS_ISOLATED)).",
1125 ));
1126 }
1127 crate::mode::validate_tool_filter_list(
1128 "available_tools",
1129 config.available_tools.as_deref(),
1130 )?;
1131 crate::mode::validate_tool_filter_list("excluded_tools", config.excluded_tools.as_deref())?;
1132 config.system_message =
1133 crate::mode::system_message_for_mode(mode, config.system_message.take());
1134 config.memory = crate::mode::memory_for_mode(mode, config.memory.take());
1135 config.enable_experimental_mode =
1136 crate::mode::experimental_mode_for_mode(mode, config.enable_experimental_mode);
1137 if mode == crate::ClientMode::Empty {
1138 if config.enable_session_telemetry.is_none() {
1139 config.enable_session_telemetry = Some(false);
1140 }
1141 if config.skip_embedding_retrieval.is_none() {
1142 config.skip_embedding_retrieval = Some(true);
1143 }
1144 if config.enable_on_demand_instruction_discovery.is_none() {
1145 config.enable_on_demand_instruction_discovery = Some(false);
1146 }
1147 if config.enable_file_hooks.is_none() {
1148 config.enable_file_hooks = Some(false);
1149 }
1150 if config.enable_host_git_operations.is_none() {
1151 config.enable_host_git_operations = Some(false);
1152 }
1153 if config.enable_session_store.is_none() {
1154 config.enable_session_store = Some(false);
1155 }
1156 if config.enable_skills.is_none() {
1157 config.enable_skills = Some(false);
1158 }
1159 }
1160 if mode == crate::ClientMode::Empty && config.mcp_oauth_token_storage.is_none() {
1161 config.mcp_oauth_token_storage = Some("in-memory".into());
1162 }
1163 if mode == crate::ClientMode::Empty && config.embedding_cache_storage.is_none() {
1164 config.embedding_cache_storage = Some("in-memory".into());
1165 }
1166 config.custom_agents_local_only =
1167 crate::mode::resolve_custom_agents_local_only(mode, config.custom_agents_local_only);
1168 let opt_skip_custom_instructions = config.skip_custom_instructions;
1169 let opt_custom_agents_local_only = config.custom_agents_local_only;
1170 let opt_coauthor_enabled = config.coauthor_enabled;
1171 let opt_manage_schedule_enabled = config.manage_schedule_enabled;
1172 let (mut wire, mut runtime) = config.into_wire()?;
1173 wire.enable_github_telemetry_forwarding =
1174 self.inner.on_github_telemetry.is_some().then_some(true);
1175
1176 let permission_handler = crate::permission::resolve_handler(
1177 runtime.permission_handler.take(),
1178 runtime.permission_policy.take(),
1179 );
1180 let handlers = SessionHandlers {
1181 permission: permission_handler,
1182 managed_settings_enabled: has_managed_settings(
1183 wire.enable_managed_settings,
1184 wire.managed_settings.as_ref(),
1185 ),
1186 elicitation: runtime.elicitation_handler.take(),
1187 mcp_auth: runtime.mcp_auth_handler.take(),
1188 user_input: runtime.user_input_handler.take(),
1189 exit_plan_mode: runtime.exit_plan_mode_handler.take(),
1190 auto_mode_switch: runtime.auto_mode_switch_handler.take(),
1191 tools: Arc::new(std::mem::take(&mut runtime.tool_handlers)),
1192 };
1193 let hooks = runtime.hooks_handler.take();
1194 let transforms = runtime.system_message_transform.take();
1195 let tools_count = wire.tools.as_ref().map_or(0, Vec::len);
1196 let commands_count = runtime.commands.as_ref().map_or(0, Vec::len);
1197 let has_hooks = hooks.is_some();
1198 let command_handlers = build_command_handler_map(runtime.commands.as_deref());
1199 let canvas_handler = runtime.canvas_handler.take();
1200 let session_fs_provider = runtime.session_fs_provider.take();
1201 let bearer_token_providers = std::mem::take(&mut runtime.bearer_token_providers);
1202 let has_mcp_auth_handler = handlers.mcp_auth.is_some();
1203 if self.inner.session_fs_configured && session_fs_provider.is_none() {
1204 return Err(ErrorKind::Session(SessionErrorKind::SessionFsProviderRequired).into());
1205 }
1206 if self.inner.session_fs_sqlite_declared
1207 && let Some(ref provider) = session_fs_provider
1208 && provider.sqlite().is_none()
1209 {
1210 return Err(Error::with_message(
1211 ErrorKind::InvalidConfig,
1212 "SessionFs capabilities declare SQLite support but the provider \
1213 does not implement SessionFsSqliteProvider",
1214 ));
1215 }
1216
1217 let mut params = serde_json::to_value(&wire)?;
1218 let trace_ctx = self.resolve_trace_context().await;
1219 inject_trace_context(&mut params, &trace_ctx);
1220
1221 let capabilities = Arc::new(parking_lot::RwLock::new(SessionCapabilities::default()));
1222 let setup_start = Instant::now();
1223 let channels = self.register_session(&session_id);
1224 let idle_waiter = Arc::new(ParkingLotMutex::new(None));
1225 let open_canvases = Arc::new(parking_lot::RwLock::new(Vec::new()));
1226 let shutdown = CancellationToken::new();
1227 let (event_tx, _) = tokio::sync::broadcast::channel(512);
1228 let event_loop = spawn_event_loop(
1229 session_id.clone(),
1230 self.clone(),
1231 handlers,
1232 hooks,
1233 transforms,
1234 command_handlers,
1235 canvas_handler,
1236 session_fs_provider,
1237 bearer_token_providers,
1238 channels,
1239 idle_waiter.clone(),
1240 capabilities.clone(),
1241 open_canvases.clone(),
1242 event_tx.clone(),
1243 shutdown.clone(),
1244 );
1245 let mut registration =
1246 PendingSessionRegistration::new(self.clone(), session_id.clone(), shutdown.clone());
1247 tracing::debug!(
1248 elapsed_ms = setup_start.elapsed().as_millis(),
1249 session_id = %session_id,
1250 tools_count,
1251 commands_count,
1252 has_hooks,
1253 "Client::resume_session local setup complete"
1254 );
1255
1256 let rpc_start = Instant::now();
1257 let result = match self.call("session.resume", Some(params)).await {
1258 Ok(result) => result,
1259 Err(error) => {
1260 registration.cleanup(event_loop).await;
1261 return Err(error);
1262 }
1263 };
1264 tracing::debug!(
1265 elapsed_ms = rpc_start.elapsed().as_millis(),
1266 session_id = %session_id,
1267 "Client::resume_session session resume request completed successfully"
1268 );
1269
1270 let resume_result: ResumeSessionResult = match serde_json::from_value(result) {
1271 Ok(result) => result,
1272 Err(error) => {
1273 registration.cleanup(event_loop).await;
1274 return Err(error.into());
1275 }
1276 };
1277 let cli_session_id = resume_result
1278 .session_id
1279 .clone()
1280 .unwrap_or_else(|| session_id.clone());
1281 if cli_session_id != session_id {
1282 registration.cleanup(event_loop).await;
1283 return Err(ErrorKind::Session(SessionErrorKind::SessionIdMismatch {
1284 requested: session_id,
1285 returned: cli_session_id,
1286 })
1287 .into());
1288 }
1289 if has_mcp_auth_handler {
1290 register_mcp_auth_interest(self, &session_id).await?;
1291 }
1292
1293 let skills_reload_start = Instant::now();
1295 if let Err(e) = self
1296 .call(
1297 "session.skills.reload",
1298 Some(serde_json::json!({ "sessionId": session_id })),
1299 )
1300 .await
1301 {
1302 warn!(
1303 elapsed_ms = skills_reload_start.elapsed().as_millis(),
1304 session_id = %session_id,
1305 error = %e,
1306 "Client::resume_session skills reload request failed"
1307 );
1308 } else {
1309 tracing::debug!(
1310 elapsed_ms = skills_reload_start.elapsed().as_millis(),
1311 session_id = %session_id,
1312 "Client::resume_session skills reload request completed successfully"
1313 );
1314 }
1315
1316 *capabilities.write() = resume_result.capabilities.unwrap_or_default();
1317 {
1322 let mut snapshots = open_canvases.write();
1323 for snapshot in resume_result.open_canvases.unwrap_or_default() {
1324 upsert_open_canvas_snapshot(&mut snapshots, snapshot);
1325 }
1326 }
1327
1328 tracing::debug!(
1329 elapsed_ms = total_start.elapsed().as_millis(),
1330 session_id = %session_id,
1331 "Client::resume_session complete"
1332 );
1333 registration.disarm();
1334 let session = Session {
1335 id: session_id,
1336 cwd: self.cwd().clone(),
1337 workspace_path: resume_result.workspace_path,
1338 remote_url: resume_result.remote_url,
1339 client: self.clone(),
1340 event_loop: ParkingLotMutex::new(Some(event_loop)),
1341 shutdown,
1342 idle_waiter,
1343 capabilities,
1344 open_canvases,
1345 event_tx,
1346 };
1347 apply_mode_post_create_patch(
1348 &session,
1349 mode,
1350 opt_skip_custom_instructions,
1351 opt_custom_agents_local_only,
1352 opt_coauthor_enabled,
1353 opt_manage_schedule_enabled,
1354 )
1355 .await?;
1356 Ok(session)
1357 }
1358}
1359
1360type CommandHandlerMap = HashMap<String, Arc<dyn CommandHandler>>;
1361
1362async fn apply_mode_post_create_patch(
1363 session: &Session,
1364 mode: crate::ClientMode,
1365 opt_skip_custom_instructions: Option<bool>,
1366 opt_custom_agents_local_only: Option<bool>,
1367 opt_coauthor_enabled: Option<bool>,
1368 opt_manage_schedule_enabled: Option<bool>,
1369) -> Result<(), Error> {
1370 use crate::generated::api_types::SessionUpdateOptionsParams;
1371 let mut patch = SessionUpdateOptionsParams::default();
1372 let should_send = if mode == crate::ClientMode::Empty {
1373 patch.skip_custom_instructions = Some(opt_skip_custom_instructions.unwrap_or(true));
1374 patch.custom_agents_local_only = Some(opt_custom_agents_local_only.unwrap_or(true));
1375 patch.coauthor_enabled = Some(opt_coauthor_enabled.unwrap_or(false));
1376 patch.manage_schedule_enabled = Some(opt_manage_schedule_enabled.unwrap_or(false));
1377 patch.installed_plugins = Some(Vec::new());
1378 true
1379 } else {
1380 let mut any = false;
1381 if let Some(v) = opt_skip_custom_instructions {
1382 patch.skip_custom_instructions = Some(v);
1383 any = true;
1384 }
1385 if let Some(v) = opt_custom_agents_local_only {
1386 patch.custom_agents_local_only = Some(v);
1387 any = true;
1388 }
1389 if let Some(v) = opt_coauthor_enabled {
1390 patch.coauthor_enabled = Some(v);
1391 any = true;
1392 }
1393 if let Some(v) = opt_manage_schedule_enabled {
1394 patch.manage_schedule_enabled = Some(v);
1395 any = true;
1396 }
1397 any
1398 };
1399 if !should_send {
1400 return Ok(());
1401 }
1402 if let Err(error) = session.rpc().options().update(patch).await {
1403 let _ = session.disconnect().await;
1404 return Err(error);
1405 }
1406 Ok(())
1407}
1408
1409fn build_command_handler_map(commands: Option<&[CommandDefinition]>) -> Arc<CommandHandlerMap> {
1410 let map = match commands {
1411 Some(commands) => commands
1412 .iter()
1413 .filter(|cmd| !cmd.name.is_empty())
1414 .map(|cmd| (cmd.name.clone(), cmd.handler.clone()))
1415 .collect(),
1416 None => HashMap::new(),
1417 };
1418 Arc::new(map)
1419}
1420
1421fn upsert_open_canvas_snapshot(
1422 snapshots: &mut Vec<OpenCanvasInstance>,
1423 snapshot: OpenCanvasInstance,
1424) {
1425 if let Some(existing) = snapshots
1426 .iter_mut()
1427 .find(|open| open.instance_id == snapshot.instance_id)
1428 {
1429 *existing = snapshot;
1430 } else {
1431 snapshots.push(snapshot);
1432 }
1433}
1434
1435fn remove_open_canvas_snapshot(snapshots: &mut Vec<OpenCanvasInstance>, instance_id: &str) {
1436 snapshots.retain(|open| open.instance_id != instance_id);
1437}
1438
1439#[allow(clippy::too_many_arguments)]
1440fn spawn_event_loop(
1441 session_id: SessionId,
1442 client: Client,
1443 handlers: SessionHandlers,
1444 hooks: Option<Arc<dyn SessionHooks>>,
1445 transforms: Option<Arc<dyn SystemMessageTransform>>,
1446 command_handlers: Arc<CommandHandlerMap>,
1447 canvas_handler: Option<Arc<dyn CanvasHandler>>,
1448 session_fs_provider: Option<Arc<dyn SessionFsProvider>>,
1449 bearer_token_providers: HashMap<String, Arc<dyn BearerTokenProvider>>,
1450 channels: crate::router::SessionChannels,
1451 idle_waiter: Arc<ParkingLotMutex<Option<IdleWaiter>>>,
1452 capabilities: Arc<parking_lot::RwLock<SessionCapabilities>>,
1453 open_canvases: Arc<parking_lot::RwLock<Vec<OpenCanvasInstance>>>,
1454 event_tx: tokio::sync::broadcast::Sender<SessionEvent>,
1455 shutdown: CancellationToken,
1456) -> JoinHandle<()> {
1457 let crate::router::SessionChannels {
1458 mut notifications,
1459 mut requests,
1460 } = channels;
1461
1462 let span = tracing::error_span!("session_event_loop", session_id = %session_id);
1463 tokio::spawn(
1464 async move {
1465 loop {
1466 tokio::select! {
1490 _ = shutdown.cancelled() => break,
1491 Some(notification) = notifications.recv() => {
1492 handle_notification(
1493 &session_id, &client, &handlers, &command_handlers, notification, &idle_waiter, &capabilities, &open_canvases, &event_tx,
1494 ).await;
1495 }
1496 Some(request) = requests.recv() => {
1497 let span = tracing::error_span!("session_request_handler", session_id = %session_id);
1501 let session_id = session_id.clone();
1502 let client = client.clone();
1503 let handlers = handlers.clone();
1504 let hooks = hooks.clone();
1505 let transforms = transforms.clone();
1506 let canvas_handler = canvas_handler.clone();
1507 let session_fs_provider = session_fs_provider.clone();
1508 let bearer_token_providers = bearer_token_providers.clone();
1509 tokio::spawn(
1510 async move {
1511 let ctx = RequestDispatchContext {
1512 client: &client,
1513 handlers: &handlers,
1514 hooks: hooks.as_deref(),
1515 transforms: transforms.as_deref(),
1516 canvas_handler: canvas_handler.as_ref(),
1517 session_fs_provider: session_fs_provider.as_ref(),
1518 bearer_token_providers: &bearer_token_providers,
1519 };
1520 handle_request(&session_id, ctx, request).await;
1521 }
1522 .instrument(span),
1523 );
1524 }
1525 else => break,
1526 }
1527 }
1528 if let Some(waiter) = idle_waiter.lock().take() {
1531 let _ = waiter
1532 .tx
1533 .send(Err(ErrorKind::Session(SessionErrorKind::EventLoopClosed).into()));
1534 }
1535 }
1536 .instrument(span),
1537 )
1538}
1539
1540fn extract_request_id(data: &Value) -> Option<RequestId> {
1541 data.get("requestId")
1542 .and_then(|v| v.as_str())
1543 .filter(|s| !s.is_empty())
1544 .map(RequestId::new)
1545}
1546
1547fn permission_request_data(
1548 event_data: &Value,
1549 managed_settings_enabled: bool,
1550) -> PermissionRequestData {
1551 let request_data = event_data
1552 .get("permissionRequest")
1553 .cloned()
1554 .unwrap_or_else(|| event_data.clone());
1555 let managed_approval_required = match request_data.get("managedApprovalRequired") {
1556 None => None,
1557 Some(Value::Bool(value)) => Some(*value),
1558 Some(_) => Some(true),
1559 };
1560 match serde_json::from_value::<PermissionRequestData>(request_data) {
1561 Ok(mut data) => {
1562 data.extra = event_data.clone();
1563 data.managed_settings_enabled = managed_settings_enabled;
1564 data
1565 }
1566 Err(_) => PermissionRequestData {
1567 kind: None,
1568 tool_call_id: None,
1569 managed_approval_required,
1570 managed_settings_enabled,
1571 extra: event_data.clone(),
1572 },
1573 }
1574}
1575
1576fn permission_response_params(
1584 session_id: &SessionId,
1585 request_id: &RequestId,
1586 result: &PermissionResult,
1587) -> Option<Value> {
1588 let (decision, decision_context) = match result {
1589 PermissionResult::Decision { decision, context } => (decision, context.clone()),
1590 PermissionResult::NoResult => return None,
1591 };
1592 let mut params = serde_json::to_value(PermissionDecisionRequest {
1593 decision_context,
1594 request_id: request_id.clone(),
1595 result: decision.clone(),
1596 })
1597 .expect("serializing permission response should succeed");
1598 params["sessionId"] =
1599 serde_json::to_value(session_id).expect("serializing session ID should succeed");
1600 Some(params)
1601}
1602
1603async fn register_mcp_auth_interest(client: &Client, session_id: &SessionId) -> Result<(), Error> {
1604 let mut params = serde_json::to_value(RegisterEventInterestParams {
1605 event_type: "mcp.oauth_required".to_string(),
1606 })?;
1607 params["sessionId"] = Value::String(session_id.to_string());
1608 client
1609 .call(rpc_methods::SESSION_EVENTLOG_REGISTERINTEREST, Some(params))
1610 .await?;
1611 Ok(())
1612}
1613
1614fn tool_failure_result(message: impl Into<String>) -> ToolResult {
1615 let message = message.into();
1616 ToolResult::Expanded(ToolResultExpanded {
1617 text_result_for_llm: message.clone(),
1618 result_type: "failure".to_string(),
1619 binary_results_for_llm: None,
1620 session_log: None,
1621 error: Some(message),
1622 tool_telemetry: None,
1623 tool_references: None,
1624 })
1625}
1626
1627#[allow(clippy::too_many_arguments)]
1629async fn handle_notification(
1630 session_id: &SessionId,
1631 client: &Client,
1632 handlers: &SessionHandlers,
1633 command_handlers: &Arc<CommandHandlerMap>,
1634 notification: SessionEventNotification,
1635 idle_waiter: &Arc<ParkingLotMutex<Option<IdleWaiter>>>,
1636 capabilities: &Arc<parking_lot::RwLock<SessionCapabilities>>,
1637 open_canvases: &Arc<parking_lot::RwLock<Vec<OpenCanvasInstance>>>,
1638 event_tx: &tokio::sync::broadcast::Sender<SessionEvent>,
1639) {
1640 let dispatch_start = Instant::now();
1641 let event = notification.event.clone();
1642 let event_type = event.parsed_type();
1643 if event_type == SessionEventType::PermissionRequested {
1644 tracing::debug!(
1645 session_id = %session_id,
1646 event_type = %event.event_type,
1647 "Session::handle_notification permission request received"
1648 );
1649 }
1650
1651 match event_type {
1654 SessionEventType::AssistantMessage
1655 | SessionEventType::SessionIdle
1656 | SessionEventType::SessionError => {
1657 let mut guard = idle_waiter.lock();
1658 if let Some(waiter) = guard.as_mut() {
1659 match event_type {
1660 SessionEventType::AssistantMessage => {
1661 if !waiter.first_assistant_message_seen {
1662 waiter.first_assistant_message_seen = true;
1663 tracing::debug!(
1664 elapsed_ms = waiter.started_at.elapsed().as_millis(),
1665 session_id = %session_id,
1666 "Session::send_and_wait first assistant message"
1667 );
1668 }
1669 waiter.last_assistant_message = Some(event.clone());
1670 }
1671 SessionEventType::SessionIdle | SessionEventType::SessionError => {
1672 if let Some(waiter) = guard.take() {
1673 if event_type == SessionEventType::SessionIdle {
1674 tracing::debug!(
1675 elapsed_ms = waiter.started_at.elapsed().as_millis(),
1676 session_id = %session_id,
1677 "Session::send_and_wait idle received"
1678 );
1679 let _ = waiter.tx.send(Ok(waiter.last_assistant_message));
1680 } else {
1681 let error_msg = event
1682 .typed_data::<SessionErrorData>()
1683 .map(|d| d.message)
1684 .or_else(|| {
1685 event
1686 .data
1687 .get("message")
1688 .and_then(|v| v.as_str())
1689 .map(|s| s.to_string())
1690 })
1691 .unwrap_or_else(|| "session error".to_string());
1692 let _ = waiter.tx.send(Err(Error::with_message(
1693 ErrorKind::Session(SessionErrorKind::AgentError),
1694 error_msg,
1695 )));
1696 }
1697 }
1698 }
1699 _ => {}
1700 }
1701 }
1702 }
1703 _ => {}
1704 }
1705
1706 if event_type == SessionEventType::CapabilitiesChanged {
1710 match serde_json::from_value::<SessionCapabilities>(notification.event.data.clone()) {
1711 Ok(changed) => *capabilities.write() = changed,
1712 Err(e) => warn!(error = %e, "failed to deserialize capabilities.changed payload"),
1713 }
1714 }
1715 if event_type == SessionEventType::SessionCanvasOpened {
1716 match serde_json::from_value::<OpenCanvasInstance>(notification.event.data.clone()) {
1717 Ok(open_canvas) => {
1718 upsert_open_canvas_snapshot(&mut open_canvases.write(), open_canvas);
1719 }
1720 Err(e) => warn!(error = %e, "failed to deserialize session.canvas.opened payload"),
1721 }
1722 }
1723 if event_type == SessionEventType::SessionCanvasClosed {
1724 match serde_json::from_value::<SessionCanvasClosedData>(notification.event.data.clone()) {
1725 Ok(closed) => {
1726 if closed.instance_id.is_empty() {
1727 warn!("failed to deserialize session.canvas.closed payload");
1728 } else {
1729 remove_open_canvas_snapshot(&mut open_canvases.write(), &closed.instance_id);
1730 }
1731 }
1732 Err(e) => warn!(error = %e, "failed to deserialize session.canvas.closed payload"),
1733 }
1734 }
1735
1736 let _ = event_tx.send(event.clone());
1740
1741 tracing::debug!(
1742 elapsed_ms = dispatch_start.elapsed().as_millis(),
1743 session_id = %session_id,
1744 event_type = %notification.event.event_type,
1745 "Session::handle_notification dispatch"
1746 );
1747
1748 match event_type {
1751 SessionEventType::PermissionRequested => {
1752 let Some(request_id) = extract_request_id(¬ification.event.data) else {
1753 return;
1754 };
1755 if notification
1759 .event
1760 .data
1761 .get("resolvedByHook")
1762 .and_then(|v| v.as_bool())
1763 .unwrap_or(false)
1764 {
1765 return;
1766 }
1767 let Some(permission_handler) = handlers.permission.clone() else {
1771 return;
1772 };
1773 let client = client.clone();
1774 let sid = session_id.clone();
1775 let data = permission_request_data(
1776 ¬ification.event.data,
1777 handlers.managed_settings_enabled,
1778 );
1779 let span = tracing::error_span!(
1780 "permission_request_handler",
1781 session_id = %sid,
1782 request_id = %request_id
1783 );
1784 tokio::spawn(
1785 async move {
1786 let handler_start = Instant::now();
1787 let result = permission_handler
1788 .handle(sid.clone(), request_id.clone(), data)
1789 .await;
1790 tracing::debug!(
1791 elapsed_ms = handler_start.elapsed().as_millis(),
1792 session_id = %sid,
1793 request_id = %request_id,
1794 "PermissionHandler::handle dispatch"
1795 );
1796 let Some(params) = permission_response_params(&sid, &request_id, &result)
1797 else {
1798 return;
1802 };
1803 let rpc_start = Instant::now();
1804 let _ = client
1805 .call(
1806 rpc_methods::SESSION_PERMISSIONS_HANDLEPENDINGPERMISSIONREQUEST,
1807 Some(params),
1808 )
1809 .await;
1810 tracing::debug!(
1811 elapsed_ms = rpc_start.elapsed().as_millis(),
1812 session_id = %sid,
1813 request_id = %request_id,
1814 "Session::handle_notification response sent successfully"
1815 );
1816 }
1817 .instrument(span),
1818 );
1819 }
1820 SessionEventType::ExternalToolRequested => {
1821 let Some(request_id) = extract_request_id(¬ification.event.data) else {
1822 return;
1823 };
1824 let data: ExternalToolRequestedData =
1825 match serde_json::from_value(notification.event.data.clone()) {
1826 Ok(d) => d,
1827 Err(e) => {
1828 warn!(error = %e, "failed to deserialize external_tool.requested");
1829 let client = client.clone();
1830 let sid = session_id.clone();
1831 let span = tracing::error_span!(
1832 "external_tool_deserialize_error",
1833 session_id = %sid,
1834 request_id = %request_id
1835 );
1836 tokio::spawn(
1837 async move {
1838 let rpc_start = Instant::now();
1839 let _ = client
1840 .call(
1841 "session.tools.handlePendingToolCall",
1842 Some(serde_json::json!({
1843 "sessionId": sid,
1844 "requestId": request_id,
1845 "error": format!("Failed to deserialize tool request: {e}"),
1846 })),
1847 )
1848 .await;
1849 tracing::debug!(
1850 elapsed_ms = rpc_start.elapsed().as_millis(),
1851 session_id = %sid,
1852 request_id = %request_id,
1853 "Session::handle_notification response sent successfully"
1854 );
1855 }
1856 .instrument(span),
1857 );
1858 return;
1859 }
1860 };
1861 let tool_handler = if data.tool_name.is_empty() {
1865 None
1866 } else {
1867 handlers.tools.get(&data.tool_name).cloned()
1868 };
1869 let Some(tool_handler) = tool_handler else {
1870 return;
1871 };
1872 let client = client.clone();
1873 let sid = session_id.clone();
1874 let span = tracing::error_span!(
1875 "external_tool_handler",
1876 session_id = %sid,
1877 request_id = %request_id
1878 );
1879 tokio::spawn(
1880 async move {
1881 if data.tool_call_id.is_empty() {
1886 let error_msg = "Missing toolCallId";
1887 let rpc_start = Instant::now();
1888 let _ = client
1889 .call(
1890 "session.tools.handlePendingToolCall",
1891 Some(serde_json::json!({
1892 "sessionId": sid,
1893 "requestId": request_id,
1894 "error": error_msg,
1895 })),
1896 )
1897 .await;
1898 tracing::debug!(
1899 elapsed_ms = rpc_start.elapsed().as_millis(),
1900 session_id = %sid,
1901 request_id = %request_id,
1902 "Session::handle_notification response sent successfully"
1903 );
1904 return;
1905 }
1906 let tool_call_id = data.tool_call_id.clone();
1907 let tool_name = data.tool_name.clone();
1908 let available_tools = if tool_name == TOOL_SEARCH_TOOL_NAME {
1915 match client
1916 .call(
1917 rpc_methods::SESSION_TOOLS_GETCURRENTMETADATA,
1918 Some(serde_json::json!({ "sessionId": sid })),
1919 )
1920 .await
1921 {
1922 Ok(value) => {
1923 serde_json::from_value::<ToolsGetCurrentMetadataResult>(value)
1924 .ok()
1925 .and_then(|result| result.tools)
1926 }
1927 Err(_) => None,
1928 }
1929 } else {
1930 None
1931 };
1932 let invocation = ToolInvocation {
1933 session_id: sid.clone(),
1934 tool_call_id: data.tool_call_id,
1935 tool_name: data.tool_name,
1936 arguments: data
1937 .arguments
1938 .unwrap_or(Value::Object(serde_json::Map::new())),
1939 available_tools,
1940 traceparent: data.traceparent,
1941 tracestate: data.tracestate,
1942 };
1943 let handler_start = Instant::now();
1944 let tool_result = match tool_handler.call(invocation).await {
1945 Ok(r) => r,
1946 Err(e) => tool_failure_result(e.to_string()),
1947 };
1948 tracing::debug!(
1949 elapsed_ms = handler_start.elapsed().as_millis(),
1950 session_id = %sid,
1951 request_id = %request_id,
1952 tool_call_id = %tool_call_id,
1953 tool_name = %tool_name,
1954 "ToolHandler::call dispatch"
1955 );
1956 let result_value = serde_json::to_value(tool_result).unwrap_or(Value::Null);
1957 let rpc_start = Instant::now();
1958 let _ = client
1959 .call(
1960 "session.tools.handlePendingToolCall",
1961 Some(serde_json::json!({
1962 "sessionId": sid,
1963 "requestId": request_id,
1964 "result": result_value,
1965 })),
1966 )
1967 .await;
1968 tracing::debug!(
1969 elapsed_ms = rpc_start.elapsed().as_millis(),
1970 session_id = %sid,
1971 request_id = %request_id,
1972 tool_call_id = %tool_call_id,
1973 tool_name = %tool_name,
1974 "Session::handle_notification response sent successfully"
1975 );
1976 }
1977 .instrument(span),
1978 );
1979 }
1980 SessionEventType::UserInputRequested => {
1981 }
1988 SessionEventType::ElicitationRequested => {
1989 let Some(request_id) = extract_request_id(¬ification.event.data) else {
1990 return;
1991 };
1992 let Some(elicitation_handler) = handlers.elicitation.clone() else {
1996 return;
1997 };
1998 let elicitation_data: ElicitationRequestedData =
1999 match serde_json::from_value(notification.event.data.clone()) {
2000 Ok(d) => d,
2001 Err(e) => {
2002 warn!(error = %e, "failed to deserialize elicitation request");
2003 return;
2004 }
2005 };
2006 let request = ElicitationRequest {
2007 message: elicitation_data.message,
2008 requested_schema: elicitation_data
2009 .requested_schema
2010 .map(|s| serde_json::to_value(s).unwrap_or(Value::Null)),
2011 mode: elicitation_data.mode.map(|m| match m {
2012 crate::generated::session_events::ElicitationRequestedMode::Form => {
2013 crate::types::ElicitationMode::Form
2014 }
2015 crate::generated::session_events::ElicitationRequestedMode::Url => {
2016 crate::types::ElicitationMode::Url
2017 }
2018 _ => crate::types::ElicitationMode::Unknown,
2019 }),
2020 elicitation_source: elicitation_data.elicitation_source,
2021 url: elicitation_data.url,
2022 };
2023 let client = client.clone();
2024 let sid = session_id.clone();
2025 let span = tracing::error_span!(
2026 "elicitation_request_handler",
2027 session_id = %sid,
2028 request_id = %request_id
2029 );
2030 tokio::spawn(
2031 async move {
2032 let cancel = ElicitationResult {
2033 action: "cancel".to_string(),
2034 content: None,
2035 };
2036 let handler_task = tokio::spawn({
2038 let sid = sid.clone();
2039 let request_id = request_id.clone();
2040 let span = tracing::error_span!(
2041 "elicitation_callback",
2042 session_id = %sid,
2043 request_id = %request_id
2044 );
2045 async move {
2046 let handler_start = Instant::now();
2047 let response = elicitation_handler
2048 .handle(sid.clone(), request_id.clone(), request)
2049 .await;
2050 tracing::debug!(
2051 elapsed_ms = handler_start.elapsed().as_millis(),
2052 session_id = %sid,
2053 request_id = %request_id,
2054 "ElicitationHandler::handle dispatch"
2055 );
2056 response
2057 }
2058 .instrument(span)
2059 });
2060 let result = match handler_task.await {
2061 Ok(r) => r,
2062 Err(_) => cancel.clone(),
2063 };
2064 let rpc_start = Instant::now();
2065 if let Err(e) = client
2066 .call(
2067 "session.ui.handlePendingElicitation",
2068 Some(serde_json::json!({
2069 "sessionId": sid,
2070 "requestId": request_id,
2071 "result": result,
2072 })),
2073 )
2074 .await
2075 {
2076 warn!(error = %e, "handlePendingElicitation failed, sending cancel");
2078 let _ = client
2079 .call(
2080 "session.ui.handlePendingElicitation",
2081 Some(serde_json::json!({
2082 "sessionId": sid,
2083 "requestId": request_id,
2084 "result": cancel,
2085 })),
2086 )
2087 .await;
2088 } else {
2089 tracing::debug!(
2090 elapsed_ms = rpc_start.elapsed().as_millis(),
2091 session_id = %sid,
2092 request_id = %request_id,
2093 "Session::handle_notification response sent successfully"
2094 );
2095 }
2096 }
2097 .instrument(span),
2098 );
2099 }
2100 SessionEventType::McpOauthRequired => {
2101 let Some(request_id) = extract_request_id(¬ification.event.data) else {
2102 return;
2103 };
2104 let Some(mcp_auth_handler) = handlers.mcp_auth.clone() else {
2105 warn!(
2106 session_id = %session_id,
2107 request_id = %request_id,
2108 "received MCP OAuth request without a registered MCP auth handler"
2109 );
2110 return;
2111 };
2112 let data: McpOauthRequiredData =
2113 match serde_json::from_value(notification.event.data.clone()) {
2114 Ok(d) => d,
2115 Err(e) => {
2116 warn!(error = %e, "failed to deserialize MCP OAuth request");
2117 return;
2118 }
2119 };
2120 let request = McpAuthRequest {
2121 request_id: request_id.clone(),
2122 server_name: data.server_name,
2123 server_url: data.server_url,
2124 reason: data.reason,
2125 www_authenticate_params: data.www_authenticate_params,
2126 resource_metadata: data.resource_metadata,
2127 static_client_config: data.static_client_config,
2128 };
2129 let client = client.clone();
2130 let sid = session_id.clone();
2131 let span = tracing::error_span!(
2132 "mcp_auth_request_handler",
2133 session_id = %sid,
2134 request_id = %request_id
2135 );
2136 tokio::spawn(
2137 async move {
2138 let cancel = McpAuthResult::Cancelled;
2139 let handler_task = tokio::spawn({
2140 let sid = sid.clone();
2141 let request_id = request_id.clone();
2142 let span = tracing::error_span!(
2143 "mcp_auth_callback",
2144 session_id = %sid,
2145 request_id = %request_id
2146 );
2147 async move {
2148 let handler_start = Instant::now();
2149 let response = mcp_auth_handler
2150 .handle(sid.clone(), request_id.clone(), request)
2151 .await;
2152 tracing::debug!(
2153 elapsed_ms = handler_start.elapsed().as_millis(),
2154 session_id = %sid,
2155 request_id = %request_id,
2156 "McpAuthHandler::handle dispatch"
2157 );
2158 response
2159 }
2160 .instrument(span)
2161 });
2162 let result = match handler_task.await {
2163 Ok(result) => result,
2164 Err(_) => cancel,
2165 };
2166 let rpc_start = Instant::now();
2167 let _ = client
2168 .call(
2169 "session.mcp.oauth.handlePendingRequest",
2170 Some(serde_json::json!({
2171 "sessionId": sid,
2172 "requestId": request_id,
2173 "result": result.into_wire(),
2174 })),
2175 )
2176 .await;
2177 tracing::debug!(
2178 elapsed_ms = rpc_start.elapsed().as_millis(),
2179 "Session::handle_notification MCP auth response sent"
2180 );
2181 }
2182 .instrument(span),
2183 );
2184 }
2185 SessionEventType::CommandExecute => {
2186 let data: CommandExecuteData =
2187 match serde_json::from_value(notification.event.data.clone()) {
2188 Ok(d) => d,
2189 Err(e) => {
2190 warn!(error = %e, "failed to deserialize command.execute");
2191 return;
2192 }
2193 };
2194 let client = client.clone();
2195 let command_handlers = command_handlers.clone();
2196 let sid = session_id.clone();
2197 let span = tracing::error_span!("command_handler", session_id = %sid);
2198 tokio::spawn(
2199 async move {
2200 let request_id = data.request_id;
2201 let ack_error = match command_handlers.get(&data.command_name).cloned() {
2202 None => Some(format!("Unknown command: {}", data.command_name)),
2203 Some(handler) => {
2204 let command_name = data.command_name.clone();
2205 let ctx = CommandContext {
2206 session_id: sid.clone(),
2207 command: data.command,
2208 command_name: data.command_name,
2209 args: data.args,
2210 };
2211 let handler_start = Instant::now();
2212 let result = handler.on_command(ctx).await;
2213 tracing::debug!(
2214 elapsed_ms = handler_start.elapsed().as_millis(),
2215 session_id = %sid,
2216 request_id = %request_id,
2217 command_name = %command_name,
2218 "CommandHandler::call dispatch"
2219 );
2220 match result {
2221 Ok(()) => None,
2222 Err(e) => Some(e.to_string()),
2223 }
2224 }
2225 };
2226 let mut params = serde_json::json!({
2227 "sessionId": sid,
2228 "requestId": request_id,
2229 });
2230 if let Some(error_msg) = ack_error {
2231 params["error"] = serde_json::Value::String(error_msg);
2232 }
2233 let rpc_start = Instant::now();
2234 let _ = client
2235 .call("session.commands.handlePendingCommand", Some(params))
2236 .await;
2237 tracing::debug!(
2238 elapsed_ms = rpc_start.elapsed().as_millis(),
2239 session_id = %sid,
2240 request_id = %request_id,
2241 "Session::handle_notification response sent successfully"
2242 );
2243 }
2244 .instrument(span),
2245 );
2246 }
2247 _ => {}
2248 }
2249}
2250
2251struct RequestDispatchContext<'a> {
2252 client: &'a Client,
2253 handlers: &'a SessionHandlers,
2254 hooks: Option<&'a dyn SessionHooks>,
2255 transforms: Option<&'a dyn SystemMessageTransform>,
2256 canvas_handler: Option<&'a Arc<dyn CanvasHandler>>,
2257 session_fs_provider: Option<&'a Arc<dyn SessionFsProvider>>,
2258 bearer_token_providers: &'a HashMap<String, Arc<dyn BearerTokenProvider>>,
2259}
2260
2261async fn handle_request(
2263 session_id: &SessionId,
2264 ctx: RequestDispatchContext<'_>,
2265 request: crate::JsonRpcRequest,
2266) {
2267 let sid = session_id.clone();
2268 let client = ctx.client;
2269 let handlers = ctx.handlers;
2270 let hooks = ctx.hooks;
2271 let transforms = ctx.transforms;
2272 let canvas_handler = ctx.canvas_handler;
2273 let session_fs_provider = ctx.session_fs_provider;
2274 let bearer_token_providers = ctx.bearer_token_providers;
2275
2276 if request.method.starts_with("sessionFs.") {
2277 crate::session_fs_dispatch::dispatch(client, session_fs_provider, request).await;
2278 return;
2279 }
2280
2281 if request.method.starts_with("canvas.") {
2282 crate::canvas_dispatch::dispatch(client, canvas_handler, request).await;
2283 return;
2284 }
2285
2286 if request.method == crate::generated::api_types::rpc_methods::PROVIDERTOKEN_GETTOKEN {
2287 crate::provider_token_dispatch::dispatch(client, bearer_token_providers, request).await;
2288 return;
2289 }
2290
2291 match request.method.as_str() {
2292 "hooks.invoke" => {
2293 let params = request.params.as_ref();
2294 let hook_type = params
2295 .and_then(|p| p.get("hookType"))
2296 .and_then(|v| v.as_str())
2297 .unwrap_or("");
2298 let input = params
2299 .and_then(|p| p.get("input"))
2300 .cloned()
2301 .unwrap_or(Value::Object(Default::default()));
2302
2303 let rpc_result = if let Some(hooks) = hooks {
2304 match crate::hooks::dispatch_hook(hooks, &sid, hook_type, input).await {
2305 Ok(output) => output,
2306 Err(e) => {
2307 warn!(error = %e, hook_type = hook_type, "hook dispatch failed");
2308 serde_json::json!({ "output": {} })
2309 }
2310 }
2311 } else {
2312 serde_json::json!({ "output": {} })
2313 };
2314
2315 let rpc_response = JsonRpcResponse {
2316 jsonrpc: "2.0".to_string(),
2317 id: request.id,
2318 result: Some(rpc_result),
2319 error: None,
2320 };
2321 let _ = client.send_response(&rpc_response).await;
2322 }
2323
2324 "userInput.request" => {
2325 let params = request.params.as_ref();
2326 let Some(question) = params
2327 .and_then(|p| p.get("question"))
2328 .and_then(|v| v.as_str())
2329 else {
2330 warn!("userInput.request missing 'question' field");
2331 let rpc_response = JsonRpcResponse {
2332 jsonrpc: "2.0".to_string(),
2333 id: request.id,
2334 result: None,
2335 error: Some(crate::JsonRpcError {
2336 code: error_codes::INVALID_PARAMS,
2337 message: "missing required field: question".to_string(),
2338 data: None,
2339 }),
2340 };
2341 let _ = client.send_response(&rpc_response).await;
2342 return;
2343 };
2344 let question = question.to_string();
2345 let choices = params
2346 .and_then(|p| p.get("choices"))
2347 .and_then(|v| v.as_array())
2348 .map(|arr| {
2349 arr.iter()
2350 .filter_map(|v| v.as_str().map(|s| s.to_string()))
2351 .collect()
2352 });
2353 let allow_freeform = params
2354 .and_then(|p| p.get("allowFreeform"))
2355 .and_then(|v| v.as_bool());
2356
2357 let handler_start = Instant::now();
2358 let response = if let Some(user_input_handler) = handlers.user_input.as_ref() {
2359 user_input_handler
2360 .handle(sid.clone(), question, choices, allow_freeform)
2361 .await
2362 } else {
2363 None
2364 };
2365 tracing::debug!(
2366 elapsed_ms = handler_start.elapsed().as_millis(),
2367 session_id = %sid,
2368 "UserInputHandler::handle dispatch"
2369 );
2370
2371 let rpc_result = match response {
2372 Some(UserInputResponse {
2373 answer,
2374 was_freeform,
2375 }) => serde_json::json!({
2376 "answer": answer,
2377 "wasFreeform": was_freeform,
2378 }),
2379 None => serde_json::json!({ "noResponse": true }),
2380 };
2381 let rpc_response = JsonRpcResponse {
2382 jsonrpc: "2.0".to_string(),
2383 id: request.id,
2384 result: Some(rpc_result),
2385 error: None,
2386 };
2387 let _ = client.send_response(&rpc_response).await;
2388 }
2389
2390 "exitPlanMode.request" => {
2391 let params = request
2392 .params
2393 .as_ref()
2394 .cloned()
2395 .unwrap_or(Value::Object(serde_json::Map::new()));
2396 let data: ExitPlanModeData = match serde_json::from_value(params) {
2397 Ok(d) => d,
2398 Err(e) => {
2399 warn!(error = %e, "failed to deserialize exitPlanMode.request params, using defaults");
2400 ExitPlanModeData::default()
2401 }
2402 };
2403
2404 let rpc_result = if let Some(exit_plan_handler) = handlers.exit_plan_mode.as_ref() {
2405 let result = exit_plan_handler.handle(sid, data).await;
2406 serde_json::to_value(result).expect("ExitPlanModeResult serialization cannot fail")
2407 } else {
2408 serde_json::json!({ "approved": true })
2409 };
2410 let rpc_response = JsonRpcResponse {
2411 jsonrpc: "2.0".to_string(),
2412 id: request.id,
2413 result: Some(rpc_result),
2414 error: None,
2415 };
2416 let _ = client.send_response(&rpc_response).await;
2417 }
2418
2419 "autoModeSwitch.request" => {
2420 let error_code = request
2421 .params
2422 .as_ref()
2423 .and_then(|p| p.get("errorCode"))
2424 .and_then(|v| v.as_str())
2425 .map(|s| s.to_string());
2426 let retry_after_seconds = request
2427 .params
2428 .as_ref()
2429 .and_then(|p| p.get("retryAfterSeconds"))
2430 .and_then(|v| v.as_f64());
2431
2432 let answer = if let Some(auto_mode_handler) = handlers.auto_mode_switch.as_ref() {
2433 auto_mode_handler
2434 .handle(sid, error_code, retry_after_seconds)
2435 .await
2436 } else {
2437 AutoModeSwitchResponse::No
2438 };
2439 let rpc_response = JsonRpcResponse {
2440 jsonrpc: "2.0".to_string(),
2441 id: request.id,
2442 result: Some(serde_json::json!({ "response": answer })),
2443 error: None,
2444 };
2445 let _ = client.send_response(&rpc_response).await;
2446 }
2447
2448 "systemMessage.transform" => {
2449 let params = request.params.as_ref();
2450 let sections: HashMap<String, crate::transforms::TransformSection> =
2451 match params.and_then(|p| p.get("sections")) {
2452 Some(v) => match serde_json::from_value(v.clone()) {
2453 Ok(s) => s,
2454 Err(e) => {
2455 let _ = send_error_response(
2456 client,
2457 request.id,
2458 error_codes::INVALID_PARAMS,
2459 &format!("invalid sections: {e}"),
2460 )
2461 .await;
2462 return;
2463 }
2464 },
2465 None => {
2466 let _ = send_error_response(
2467 client,
2468 request.id,
2469 error_codes::INVALID_PARAMS,
2470 "missing sections parameter",
2471 )
2472 .await;
2473 return;
2474 }
2475 };
2476
2477 let rpc_result = if let Some(transforms) = transforms {
2478 let transform_start = Instant::now();
2479 let response =
2480 crate::transforms::dispatch_transform(transforms, &sid, sections).await;
2481 tracing::debug!(
2482 elapsed_ms = transform_start.elapsed().as_millis(),
2483 session_id = %sid,
2484 "SystemMessageTransform::transform_section dispatch"
2485 );
2486 match serde_json::to_value(response) {
2487 Ok(v) => v,
2488 Err(e) => {
2489 warn!(error = %e, "failed to serialize transform response");
2490 serde_json::json!({ "sections": {} })
2491 }
2492 }
2493 } else {
2494 let passthrough: HashMap<String, crate::transforms::TransformSection> = sections;
2496 serde_json::json!({ "sections": passthrough })
2497 };
2498
2499 let rpc_response = JsonRpcResponse {
2500 jsonrpc: "2.0".to_string(),
2501 id: request.id,
2502 result: Some(rpc_result),
2503 error: None,
2504 };
2505 let _ = client.send_response(&rpc_response).await;
2506 }
2507
2508 method => {
2509 warn!(
2510 method = method,
2511 "unhandled request method in session event loop"
2512 );
2513 let _ = send_error_response(
2514 client,
2515 request.id,
2516 error_codes::METHOD_NOT_FOUND,
2517 &format!("unknown method: {method}"),
2518 )
2519 .await;
2520 }
2521 }
2522}
2523
2524async fn send_error_response(
2525 client: &Client,
2526 id: u64,
2527 code: i32,
2528 message: &str,
2529) -> Result<(), Error> {
2530 let response = JsonRpcResponse {
2531 jsonrpc: "2.0".to_string(),
2532 id,
2533 result: None,
2534 error: Some(crate::JsonRpcError {
2535 code,
2536 message: message.to_string(),
2537 data: None,
2538 }),
2539 };
2540 client.send_response(&response).await
2541}
2542
2543fn apply_transform_sections(
2547 sys_msg: &mut SystemMessageConfig,
2548 transforms: &dyn SystemMessageTransform,
2549) {
2550 sys_msg.mode = Some("customize".to_string());
2551 let sections = sys_msg.sections.get_or_insert_with(HashMap::new);
2552 for id in transforms.section_ids() {
2553 sections.entry(id).or_insert_with(|| SectionOverride {
2554 action: Some("transform".to_string()),
2555 content: None,
2556 });
2557 }
2558}
2559
2560fn inject_transform_sections(config: &mut SessionConfig, transforms: &dyn SystemMessageTransform) {
2561 let sys_msg = config.system_message.get_or_insert_with(Default::default);
2562 apply_transform_sections(sys_msg, transforms);
2563}
2564
2565fn inject_transform_sections_resume(
2566 config: &mut ResumeSessionConfig,
2567 transforms: &dyn SystemMessageTransform,
2568) {
2569 let sys_msg = config.system_message.get_or_insert_with(Default::default);
2570 apply_transform_sections(sys_msg, transforms);
2571}
2572
2573#[cfg(test)]
2574mod tests {
2575 use serde_json::json;
2576
2577 use super::{has_managed_settings, permission_request_data, permission_response_params};
2578 use crate::handler::PermissionResult;
2579 use crate::types::{
2580 PermissionDecisionContext, PermissionDecisionOutcome, PermissionDecisionSource,
2581 PermissionDecisionSurface, RequestId, SessionId,
2582 };
2583
2584 #[test]
2585 fn direct_injection_enables_managed_safeguards() {
2586 let settings = crate::types::ManagedSettings::default();
2587 assert!(has_managed_settings(None, Some(&settings)));
2588 assert!(!has_managed_settings(None, None));
2589 }
2590
2591 fn attribution_context() -> PermissionDecisionContext {
2592 PermissionDecisionContext {
2593 outcome: PermissionDecisionOutcome::AutoApproved,
2594 source: PermissionDecisionSource::JudgeRecommendation,
2595 surface: PermissionDecisionSurface::CopilotApp,
2596 }
2597 }
2598
2599 #[test]
2600 fn response_params_omit_decision_context_without_attribution() {
2601 for (result, expected) in [
2602 (
2603 PermissionResult::approve_once(),
2604 json!({ "kind": "approve-once" }),
2605 ),
2606 (PermissionResult::reject(None), json!({ "kind": "reject" })),
2607 (
2608 PermissionResult::reject(Some("bad".to_string())),
2609 json!({ "kind": "reject", "feedback": "bad" }),
2610 ),
2611 (
2612 PermissionResult::user_not_available(),
2613 json!({ "kind": "user-not-available" }),
2614 ),
2615 ] {
2616 let params = permission_response_params(
2617 &SessionId::from("session-1"),
2618 &RequestId::from("permission-1"),
2619 &result,
2620 )
2621 .unwrap();
2622 assert_eq!(
2623 params,
2624 json!({
2625 "sessionId": "session-1",
2626 "requestId": "permission-1",
2627 "result": expected,
2628 })
2629 );
2630 }
2631 }
2632
2633 #[test]
2634 fn response_params_forward_decision_context_alongside_result() {
2635 let params = permission_response_params(
2636 &SessionId::from("session-1"),
2637 &RequestId::from("permission-1"),
2638 &PermissionResult::approve_once().with_context(attribution_context()),
2639 )
2640 .unwrap();
2641 assert_eq!(
2642 params,
2643 json!({
2644 "sessionId": "session-1",
2645 "requestId": "permission-1",
2646 "result": { "kind": "approve-once" },
2647 "decisionContext": {
2648 "outcome": "auto_approved",
2649 "source": "judge_recommendation",
2650 "surface": "copilot_app",
2651 },
2652 })
2653 );
2654 assert!(params["result"].get("decisionContext").is_none());
2656 }
2657
2658 #[test]
2659 fn response_params_suppressed_for_no_result() {
2660 assert!(
2661 permission_response_params(
2662 &SessionId::from("session-1"),
2663 &RequestId::from("permission-1"),
2664 &PermissionResult::NoResult,
2665 )
2666 .is_none()
2667 );
2668 }
2669
2670 #[test]
2671 fn with_context_is_a_no_op_on_no_result() {
2672 let result = PermissionResult::no_result().with_context(attribution_context());
2673 assert!(matches!(result, PermissionResult::NoResult));
2674 }
2675
2676 #[test]
2677 fn with_context_replaces_rather_than_nests() {
2678 let result = PermissionResult::approve_once()
2679 .with_context(attribution_context())
2680 .with_context(PermissionDecisionContext {
2681 outcome: PermissionDecisionOutcome::PromptedUser,
2682 source: PermissionDecisionSource::HumanResponse,
2683 surface: PermissionDecisionSurface::Sdk,
2684 });
2685 let params = permission_response_params(
2686 &SessionId::from("session-1"),
2687 &RequestId::from("permission-1"),
2688 &result,
2689 )
2690 .unwrap();
2691 assert_eq!(
2692 params["decisionContext"],
2693 json!({
2694 "outcome": "prompted_user",
2695 "source": "human_response",
2696 "surface": "sdk",
2697 })
2698 );
2699 }
2700
2701 #[test]
2702 fn permission_request_data_reads_nested_managed_approval_metadata() {
2703 let data = permission_request_data(
2704 &json!({
2705 "requestId": "permission-1",
2706 "permissionRequest": {
2707 "kind": "read",
2708 "managedApprovalRequired": true,
2709 "path": "/workspace/file.txt"
2710 }
2711 }),
2712 false,
2713 );
2714
2715 assert_eq!(data.managed_approval_required, Some(true));
2716 assert_eq!(
2717 data.extra["permissionRequest"]["path"],
2718 "/workspace/file.txt"
2719 );
2720 }
2721
2722 #[test]
2723 fn permission_request_data_preserves_managed_flag_when_other_fields_are_malformed() {
2724 let data = permission_request_data(
2725 &json!({
2726 "requestId": "permission-1",
2727 "permissionRequest": {
2728 "kind": "read",
2729 "managedApprovalRequired": true,
2730 "toolCallId": 42
2731 }
2732 }),
2733 false,
2734 );
2735
2736 assert_eq!(data.managed_approval_required, Some(true));
2737 assert_eq!(data.extra["requestId"], "permission-1");
2738 }
2739
2740 #[test]
2741 fn permission_request_data_fails_closed_for_malformed_managed_flag() {
2742 let data = permission_request_data(
2743 &json!({
2744 "requestId": "permission-1",
2745 "permissionRequest": {
2746 "kind": "read",
2747 "managedApprovalRequired": "yes",
2748 "path": "/workspace/file.txt"
2749 }
2750 }),
2751 false,
2752 );
2753
2754 assert_eq!(data.managed_approval_required, Some(true));
2755 }
2756
2757 #[test]
2758 fn permission_request_data_preserves_valid_false_managed_flag() {
2759 let data = permission_request_data(
2760 &json!({
2761 "requestId": "permission-1",
2762 "permissionRequest": {
2763 "kind": "read",
2764 "managedApprovalRequired": false,
2765 "path": "/workspace/file.txt"
2766 }
2767 }),
2768 false,
2769 );
2770
2771 assert_eq!(data.managed_approval_required, Some(false));
2772 }
2773}