1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum GitHubTokenRequestReason {
24 Initial,
26 Refresh,
28}
29
30#[derive(Debug, Clone, PartialEq, Eq)]
32pub struct GitHubTokenProviderArgs {
33 pub host: String,
35 pub session_id: Option<crate::SessionId>,
37 pub reason: GitHubTokenRequestReason,
39}
40
41pub struct GitHubToken {
46 access_token: String,
47 expires_in_seconds: i64,
48 token_type: Option<String>,
49}
50
51impl GitHubToken {
52 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 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
87pub enum GitHubTokenProviderResult {
89 Token(GitHubToken),
91 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_trait]
106pub trait GitHubTokenProvider: Send + Sync {
107 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(®istration_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 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 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(®istration_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(®istration_id)
448 );
449 }
450}