Skip to main content

github_copilot_sdk/
extension_launch_provider.rs

1//! Connection-level extension launch profile resolution.
2
3use std::sync::{Arc, OnceLock, Weak};
4
5use async_trait::async_trait;
6use parking_lot::RwLock;
7use serde::Serialize;
8use serde_json::Value;
9use tracing::warn;
10
11pub use crate::rpc::{
12    ExtensionLaunchProfile, ExtensionLaunchProviderResolveRequest,
13    ExtensionLaunchProviderResolveResult,
14};
15use crate::{
16    Client, ClientInner, JsonRpcError, JsonRpcRequest, JsonRpcResponse, Result, error_codes,
17};
18
19pub(crate) const RESOLVE_METHOD: &str = "extensionLaunchProvider.resolve";
20const MISSING_HANDLER_MESSAGE: &str = "No extensionLaunchProvider client-global handler registered";
21
22/// Resolves process launch profiles for extension entrypoints discovered by the runtime.
23///
24/// Configure an implementation with
25/// [`ClientOptions::with_extension_launch_provider`](crate::ClientOptions::with_extension_launch_provider).
26/// The SDK registers the provider before [`Client::start`](crate::Client::start)
27/// returns, so extension resolution cannot race session creation.
28///
29/// The returned executable, arguments, and environment are forwarded unchanged.
30/// The runtime remains authoritative for its reserved `COPILOT_SDK_PATH`,
31/// `SESSION_ID`, and `COPILOT_EXTENSION_PARENT_PID` environment variables.
32#[async_trait]
33pub trait ExtensionLaunchProvider: Send + Sync + 'static {
34    /// Resolve a launch profile for one discovered extension entrypoint.
35    ///
36    /// Return a result with `launch: None` when the provider does not support
37    /// the entrypoint.
38    async fn resolve(
39        &self,
40        request: ExtensionLaunchProviderResolveRequest,
41    ) -> Result<ExtensionLaunchProviderResolveResult>;
42}
43
44pub(crate) struct ExtensionLaunchProviderDispatcher {
45    handler: RwLock<Option<Arc<dyn ExtensionLaunchProvider>>>,
46    client: OnceLock<Weak<ClientInner>>,
47}
48
49impl ExtensionLaunchProviderDispatcher {
50    pub(crate) fn new(handler: Option<Arc<dyn ExtensionLaunchProvider>>) -> Self {
51        Self {
52            handler: RwLock::new(handler),
53            client: OnceLock::new(),
54        }
55    }
56
57    pub(crate) fn set_client(&self, client: Weak<ClientInner>) {
58        let _ = self.client.set(client);
59    }
60
61    pub(crate) fn is_configured(&self) -> bool {
62        self.handler.read().is_some()
63    }
64
65    pub(crate) fn clear(&self) {
66        self.handler.write().take();
67    }
68
69    pub(crate) async fn dispatch(&self, request: JsonRpcRequest) {
70        let request_id = request.id;
71        let Some(handler) = self.handler.read().clone() else {
72            self.send_error(
73                request_id,
74                error_codes::INTERNAL_ERROR,
75                MISSING_HANDLER_MESSAGE,
76            )
77            .await;
78            return;
79        };
80
81        let params = request
82            .params
83            .unwrap_or_else(|| Value::Object(serde_json::Map::new()));
84        let request = match serde_json::from_value(params) {
85            Ok(request) => request,
86            Err(error) => {
87                self.send_error(
88                    request_id,
89                    error_codes::INVALID_PARAMS,
90                    &format!("invalid params: {error}"),
91                )
92                .await;
93                return;
94            }
95        };
96
97        match handler.resolve(request).await {
98            Ok(result) => self.respond(request_id, result).await,
99            Err(error) => {
100                self.send_error(request_id, error_codes::INTERNAL_ERROR, &error.to_string())
101                    .await;
102            }
103        }
104    }
105
106    fn client(&self) -> Option<Client> {
107        self.client
108            .get()
109            .and_then(Weak::upgrade)
110            .map(Client::from_inner)
111    }
112
113    async fn respond<T: Serialize>(&self, request_id: u64, result: T) {
114        let value = match serde_json::to_value(result) {
115            Ok(value) => value,
116            Err(error) => {
117                warn!(error = %error, "failed to serialize extension launch provider response");
118                self.send_error(
119                    request_id,
120                    error_codes::INTERNAL_ERROR,
121                    "serialization failure",
122                )
123                .await;
124                return;
125            }
126        };
127
128        let Some(client) = self.client() else {
129            return;
130        };
131        let _ = client
132            .send_response(&JsonRpcResponse {
133                jsonrpc: "2.0".to_string(),
134                id: request_id,
135                result: Some(value),
136                error: None,
137            })
138            .await;
139    }
140
141    async fn send_error(&self, request_id: u64, code: i32, message: &str) {
142        let Some(client) = self.client() else {
143            return;
144        };
145        let _ = client
146            .send_response(&JsonRpcResponse {
147                jsonrpc: "2.0".to_string(),
148                id: request_id,
149                result: None,
150                error: Some(JsonRpcError {
151                    code,
152                    message: message.to_string(),
153                    data: None,
154                }),
155            })
156            .await;
157    }
158}