Skip to main content

github_copilot_sdk/
github_token.rs

1//! Session-scoped GitHub token provider callbacks.
2
3use std::collections::HashMap;
4use std::future::Future;
5use std::panic::AssertUnwindSafe;
6use std::sync::{Arc, OnceLock, Weak};
7
8use async_trait::async_trait;
9use futures_util::FutureExt;
10use parking_lot::Mutex;
11use serde_json::Value;
12use tokio::sync::mpsc;
13use tokio::task::JoinHandle;
14
15use crate::generated::api_types::{
16    GitHubTokenAcquireReason, GitHubTokenAcquireRequest, GitHubTokenAcquireResult,
17    GitHubTokenAcquireResultCancelled, GitHubTokenAcquireResultToken,
18};
19use crate::{Client, ClientInner, JsonRpcError, JsonRpcRequest, JsonRpcResponse, error_codes};
20
21/// Why the runtime is requesting a GitHub token.
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum GitHubTokenRequestReason {
24    /// The session needs its initial token.
25    Initial,
26    /// The session needs a refreshed token.
27    Refresh,
28}
29
30/// Context supplied when the runtime needs a GitHub token for a session.
31#[derive(Debug, Clone, PartialEq, Eq)]
32pub struct GitHubTokenProviderArgs {
33    /// Effective GitHub host for which a token is required.
34    pub host: String,
35    /// Session receiving the token, when the runtime has assigned its ID.
36    pub session_id: Option<crate::SessionId>,
37    /// Whether this is the initial token acquisition or a refresh.
38    pub reason: GitHubTokenRequestReason,
39}
40
41/// A GitHub access token returned by a session token provider.
42///
43/// `expires_in_seconds` is the positive remaining lifetime when the callback
44/// completes. Production GitHub tokens typically last eight hours.
45pub struct GitHubToken {
46    access_token: String,
47    expires_in_seconds: i64,
48    token_type: Option<String>,
49}
50
51impl GitHubToken {
52    /// Construct a token response with its remaining lifetime in seconds.
53    pub fn new(access_token: impl Into<String>, expires_in_seconds: i64) -> Self {
54        Self {
55            access_token: access_token.into(),
56            expires_in_seconds,
57            token_type: None,
58        }
59    }
60
61    /// Override the OAuth token type. The runtime defaults to `bearer` when unset.
62    pub fn with_token_type(mut self, token_type: impl Into<String>) -> Self {
63        self.token_type = Some(token_type.into());
64        self
65    }
66
67    fn into_wire(self) -> GitHubTokenAcquireResultToken {
68        GitHubTokenAcquireResultToken {
69            access_token: self.access_token,
70            expires_in: self.expires_in_seconds,
71            kind: Default::default(),
72            token_type: self.token_type,
73        }
74    }
75}
76
77impl std::fmt::Debug for GitHubToken {
78    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
79        f.debug_struct("GitHubToken")
80            .field("access_token", &"<redacted>")
81            .field("expires_in_seconds", &self.expires_in_seconds)
82            .field("token_type", &self.token_type)
83            .finish()
84    }
85}
86
87/// Result of acquiring a session-scoped GitHub token.
88pub enum GitHubTokenProviderResult {
89    /// A token was acquired.
90    Token(GitHubToken),
91    /// The host cancelled acquisition.
92    Cancelled,
93}
94
95impl std::fmt::Debug for GitHubTokenProviderResult {
96    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
97        match self {
98            Self::Token(token) => f.debug_tuple("Token").field(token).finish(),
99            Self::Cancelled => f.write_str("Cancelled"),
100        }
101    }
102}
103
104/// Async callback used to acquire GitHub tokens for one session.
105#[async_trait]
106pub trait GitHubTokenProvider: Send + Sync {
107    /// Acquire a token or explicitly cancel the request.
108    ///
109    /// Initial cancellation, errors, and invalid token responses reject session
110    /// creation or resume instead of falling back to ambient authentication.
111    async fn get_token(
112        &self,
113        args: GitHubTokenProviderArgs,
114    ) -> Result<GitHubTokenProviderResult, crate::Error>;
115}
116
117#[async_trait]
118impl<F, Fut> GitHubTokenProvider for F
119where
120    F: Fn(GitHubTokenProviderArgs) -> Fut + Send + Sync,
121    Fut: Future<Output = Result<GitHubTokenProviderResult, crate::Error>> + Send,
122{
123    async fn get_token(
124        &self,
125        args: GitHubTokenProviderArgs,
126    ) -> Result<GitHubTokenProviderResult, crate::Error> {
127        (self)(args).await
128    }
129}
130
131struct ProviderRegistration {
132    provider: Arc<dyn GitHubTokenProvider>,
133    worker: Option<TokenWorker>,
134}
135
136struct TokenWorker {
137    requests: mpsc::UnboundedSender<JsonRpcRequest>,
138    task: JoinHandle<()>,
139}
140
141impl Drop for TokenWorker {
142    fn drop(&mut self) {
143        self.task.abort();
144    }
145}
146
147#[derive(Default)]
148struct RegistryState {
149    providers: HashMap<String, ProviderRegistration>,
150    session_owners: HashMap<crate::SessionId, String>,
151}
152
153pub(crate) struct GitHubTokenRegistry {
154    state: Mutex<RegistryState>,
155    client: OnceLock<Weak<ClientInner>>,
156}
157
158impl GitHubTokenRegistry {
159    pub(crate) fn new() -> Self {
160        Self {
161            state: Mutex::new(RegistryState::default()),
162            client: OnceLock::new(),
163        }
164    }
165
166    pub(crate) fn set_client(&self, client: Weak<ClientInner>) {
167        let _ = self.client.set(client);
168    }
169
170    pub(crate) fn register(&self, provider: Arc<dyn GitHubTokenProvider>) -> String {
171        let registration_id = uuid::Uuid::new_v4().to_string();
172        self.state.lock().providers.insert(
173            registration_id.clone(),
174            ProviderRegistration {
175                provider,
176                worker: None,
177            },
178        );
179        registration_id
180    }
181
182    pub(crate) fn claim(&self, registration_id: &str, session_id: crate::SessionId) {
183        let mut state = self.state.lock();
184        if let Some(previous) = state
185            .session_owners
186            .insert(session_id, registration_id.to_string())
187            && previous != registration_id
188        {
189            state.providers.remove(&previous);
190        }
191    }
192
193    pub(crate) fn unregister(&self, registration_id: &str) {
194        let mut state = self.state.lock();
195        state.providers.remove(registration_id);
196        state
197            .session_owners
198            .retain(|_, owned| owned != registration_id);
199    }
200
201    pub(crate) fn retire_session(&self, session_id: &crate::SessionId) {
202        let mut state = self.state.lock();
203        if let Some(registration_id) = state.session_owners.remove(session_id) {
204            state.providers.remove(&registration_id);
205        }
206    }
207
208    pub(crate) fn clear(&self) {
209        let mut state = self.state.lock();
210        state.providers.clear();
211        state.session_owners.clear();
212    }
213
214    pub(crate) fn dispatch(&self, request: JsonRpcRequest) {
215        let Some(client) = self.client.get().cloned() else {
216            return;
217        };
218        let registration_id = request
219            .params
220            .as_ref()
221            .and_then(|params| params.get("registrationId"))
222            .and_then(Value::as_str);
223        let mut state = self.state.lock();
224        if let Some(registration) = registration_id.and_then(|id| state.providers.get_mut(id)) {
225            let worker = registration.worker.get_or_insert_with(|| {
226                let provider = registration.provider.clone();
227                let (requests, mut rx) = mpsc::unbounded_channel();
228                let task = tokio::spawn(async move {
229                    // One worker per registration preserves callback order without
230                    // holding up requests for other providers or sessions.
231                    while let Some(request) = rx.recv().await {
232                        Self::handle_request(&client, Some(provider.as_ref()), request).await;
233                    }
234                });
235                TokenWorker { requests, task }
236            });
237            let _ = worker.requests.send(request);
238        } else {
239            // Invalid/retired registrations still receive the normal RPC error.
240            tokio::spawn(async move {
241                Self::handle_request(&client, None, request).await;
242            });
243        }
244    }
245
246    async fn handle_request(
247        client: &Weak<ClientInner>,
248        provider: Option<&dyn GitHubTokenProvider>,
249        request: JsonRpcRequest,
250    ) {
251        let params = request
252            .params
253            .clone()
254            .unwrap_or(Value::Object(serde_json::Map::new()));
255        let params: GitHubTokenAcquireRequest = match serde_json::from_value(params) {
256            Ok(params) => params,
257            Err(error) => {
258                send_error(
259                    client,
260                    request.id,
261                    error_codes::INVALID_PARAMS,
262                    &format!("invalid params: {error}"),
263                )
264                .await;
265                return;
266            }
267        };
268        let Some(provider) = provider else {
269            send_error(
270                client,
271                request.id,
272                error_codes::INTERNAL_ERROR,
273                "unknown GitHub token provider registration",
274            )
275            .await;
276            return;
277        };
278
279        let reason = match params.reason {
280            GitHubTokenAcquireReason::Initial => GitHubTokenRequestReason::Initial,
281            GitHubTokenAcquireReason::Refresh => GitHubTokenRequestReason::Refresh,
282            GitHubTokenAcquireReason::Unknown => {
283                send_error(
284                    client,
285                    request.id,
286                    error_codes::INVALID_PARAMS,
287                    "unknown GitHub token acquisition reason",
288                )
289                .await;
290                return;
291            }
292        };
293
294        let result = AssertUnwindSafe(async {
295            provider
296                .get_token(GitHubTokenProviderArgs {
297                    host: params.host,
298                    session_id: params.session_id,
299                    reason,
300                })
301                .await
302        })
303        .catch_unwind()
304        .await;
305        let result = match result {
306            Ok(result) => result,
307            Err(_) => {
308                send_error(
309                    client,
310                    request.id,
311                    error_codes::INTERNAL_ERROR,
312                    "GitHub token provider panicked",
313                )
314                .await;
315                return;
316            }
317        };
318        match result {
319            Ok(GitHubTokenProviderResult::Token(token)) => {
320                respond(
321                    client,
322                    request.id,
323                    GitHubTokenAcquireResult::Token(token.into_wire()),
324                )
325                .await;
326            }
327            Ok(GitHubTokenProviderResult::Cancelled) => {
328                respond(
329                    client,
330                    request.id,
331                    GitHubTokenAcquireResult::Cancelled(GitHubTokenAcquireResultCancelled {
332                        kind: Default::default(),
333                    }),
334                )
335                .await;
336            }
337            Err(error) => {
338                send_error(
339                    client,
340                    request.id,
341                    error_codes::INTERNAL_ERROR,
342                    &format!("GitHub token provider failed: {error}"),
343                )
344                .await;
345            }
346        }
347    }
348}
349
350pub(crate) struct GitHubTokenRegistration {
351    registry: Arc<GitHubTokenRegistry>,
352    id: String,
353}
354
355impl GitHubTokenRegistration {
356    pub(crate) fn new(registry: Arc<GitHubTokenRegistry>, id: String) -> Self {
357        Self { registry, id }
358    }
359
360    pub(crate) fn id(&self) -> &str {
361        &self.id
362    }
363
364    pub(crate) fn claim(&self, session_id: crate::SessionId) {
365        self.registry.claim(&self.id, session_id);
366    }
367}
368
369impl Drop for GitHubTokenRegistration {
370    fn drop(&mut self) {
371        self.registry.unregister(&self.id);
372    }
373}
374
375async fn respond(client: &Weak<ClientInner>, request_id: u64, result: GitHubTokenAcquireResult) {
376    match serde_json::to_value(result) {
377        Ok(result) => {
378            let Some(inner) = client.upgrade() else {
379                return;
380            };
381            let _ = Client::from_inner(inner)
382                .send_response(&JsonRpcResponse {
383                    jsonrpc: "2.0".to_string(),
384                    id: request_id,
385                    result: Some(result),
386                    error: None,
387                })
388                .await;
389        }
390        Err(_) => {
391            send_error(
392                client,
393                request_id,
394                error_codes::INTERNAL_ERROR,
395                "serialization failure",
396            )
397            .await;
398        }
399    }
400}
401
402async fn send_error(client: &Weak<ClientInner>, request_id: u64, code: i32, message: &str) {
403    let Some(inner) = client.upgrade() else {
404        return;
405    };
406    let _ = Client::from_inner(inner)
407        .send_response(&JsonRpcResponse {
408            jsonrpc: "2.0".to_string(),
409            id: request_id,
410            result: None,
411            error: Some(JsonRpcError {
412                code,
413                message: message.to_string(),
414                data: None,
415            }),
416        })
417        .await;
418}
419
420#[cfg(test)]
421mod tests {
422    use super::*;
423
424    #[test]
425    fn token_debug_is_redacted() {
426        let token = GitHubToken::new("do-not-print", 28_800);
427        assert!(!format!("{token:?}").contains("do-not-print"));
428    }
429
430    #[test]
431    fn retiring_session_removes_its_provider() {
432        let registry = GitHubTokenRegistry::new();
433        let provider = Arc::new(|_args: GitHubTokenProviderArgs| async {
434            Ok(GitHubTokenProviderResult::Cancelled)
435        });
436        let registration_id = registry.register(provider);
437        let session_id = crate::SessionId::from("session-1");
438        registry.claim(&registration_id, session_id.clone());
439
440        registry.retire_session(&session_id);
441
442        assert!(
443            !registry
444                .state
445                .lock()
446                .providers
447                .contains_key(&registration_id)
448        );
449    }
450}