diff --git a/crates/jacquard-common/src/websocket.rs b/crates/jacquard-common/src/websocket.rs index bb150920..28415b19 100644 --- a/crates/jacquard-common/src/websocket.rs +++ b/crates/jacquard-common/src/websocket.rs @@ -475,12 +475,24 @@ impl fmt::Debug for WsSink { /// WebSocket client trait #[cfg_attr(not(target_arch = "wasm32"), trait_variant::make(Send))] -pub trait WebSocketClient { +pub trait WebSocketClient: Sync { /// Error type for WebSocket operations type Error: std::error::Error + Send + Sync + 'static; /// Connect to a WebSocket endpoint fn connect(&self, url: Url) -> impl Future>; + + /// Connect to a WebSocket endpoint with custom headers + /// + /// Default implementation ignores headers and calls `connect()`. + /// Override this method to support authentication headers for subscriptions. + fn connect_with_headers( + &self, + url: Url, + _headers: Vec<(CowStr<'_>, CowStr<'_>)>, + ) -> impl Future> { + async move { self.connect(url).await } + } } /// WebSocket connection with bidirectional streams diff --git a/crates/jacquard-oauth/Cargo.toml b/crates/jacquard-oauth/Cargo.toml index b482b960..33000d77 100644 --- a/crates/jacquard-oauth/Cargo.toml +++ b/crates/jacquard-oauth/Cargo.toml @@ -49,3 +49,4 @@ default = [] loopback = ["dep:rouille"] browser-open = ["dep:webbrowser"] tracing = ["dep:tracing"] +websocket = ["jacquard-common/websocket"] diff --git a/crates/jacquard-oauth/src/client.rs b/crates/jacquard-oauth/src/client.rs index 84ef029e..439908cd 100644 --- a/crates/jacquard-oauth/src/client.rs +++ b/crates/jacquard-oauth/src/client.rs @@ -19,6 +19,11 @@ use jacquard_common::{ build_http_request, process_response, }, }; + +#[cfg(feature = "websocket")] +use jacquard_common::websocket::{WebSocketClient, WebSocketConnection}; +#[cfg(feature = "websocket")] +use jacquard_common::xrpc::XrpcSubscription; use jacquard_identity::{ JacquardResolver, resolver::{DidDocResponse, IdentityError, IdentityResolver, ResolverOptions}, @@ -279,18 +284,19 @@ where } } -pub struct OAuthSession +pub struct OAuthSession where T: OAuthResolver, S: ClientAuthStore, { pub registry: Arc>, pub client: Arc, + pub ws_client: W, pub data: RwLock>, pub options: RwLock>, } -impl OAuthSession +impl OAuthSession where T: OAuthResolver, S: ClientAuthStore, @@ -303,6 +309,28 @@ where Self { registry, client, + ws_client: (), + data: RwLock::new(data), + options: RwLock::new(CallOptions::default()), + } + } +} + +impl OAuthSession +where + T: OAuthResolver, + S: ClientAuthStore, +{ + pub fn new_with_ws( + registry: Arc>, + client: Arc, + ws_client: W, + data: ClientSessionData<'static>, + ) -> Self { + Self { + registry, + client, + ws_client, data: RwLock::new(data), options: RwLock::new(CallOptions::default()), } @@ -312,11 +340,17 @@ where Self { registry: self.registry, client: self.client, + ws_client: self.ws_client, data: self.data, options: RwLock::new(options.into_static()), } } + /// Get a reference to the WebSocket client. + pub fn ws_client(&self) -> &W { + &self.ws_client + } + pub async fn set_options(&self, options: CallOptions<'_>) { *self.options.write().await = options.into_static(); } @@ -344,7 +378,7 @@ where .map(|t| AuthorizationToken::Dpop(t.clone())) } } -impl OAuthSession +impl OAuthSession where S: ClientAuthStore + Send + Sync + 'static, T: OAuthResolver + DpopExt + Send + Sync + 'static, @@ -373,14 +407,14 @@ where T: OAuthResolver, S: ClientAuthStore, { - pub fn from_session(session: &OAuthSession) -> Self { + pub fn from_session(session: &OAuthSession) -> Self { Self { registry: session.registry.clone(), client: session.client.clone(), } } } -impl OAuthSession +impl OAuthSession where S: ClientAuthStore + Send + Sync + 'static, T: OAuthResolver + DpopExt + Send + Sync + 'static, @@ -402,10 +436,54 @@ where } } -impl HttpClient for OAuthSession +#[cfg(feature = "websocket")] +impl OAuthSession +where + S: ClientAuthStore, + T: OAuthResolver, + W: WebSocketClient, +{ + /// Subscribe to an XRPC WebSocket subscription. + /// + /// Connects to the WebSocket endpoint and threads through DPoP authentication headers. + pub async fn subscribe(&self, params: &Sub) -> Result + where + Sub: XrpcSubscription, + { + let base_uri = self.endpoint().await; + + // Build WebSocket URL + let mut ws_url = base_uri.clone(); + ws_url.set_scheme("wss").ok(); + ws_url.set_path(&format!("/xrpc/{}", Sub::NSID)); + + // Add query params + let query_params = params.query_params(); + if !query_params.is_empty() { + let query_string = serde_html_form::to_string(&query_params).unwrap_or_default(); + ws_url.set_query(Some(&query_string)); + } + + // Thread DPoP auth headers (even though tokio-tungstenite-wasm doesn't support them yet) + let token = self.access_token().await; + let auth_value = match token { + AuthorizationToken::Bearer(t) => format!("Bearer {}", t.as_ref()), + AuthorizationToken::Dpop(t) => format!("DPoP {}", t.as_ref()), + }; + let headers = vec![( + CowStr::from("Authorization"), + CowStr::from(auth_value), + )]; + + self.ws_client.connect_with_headers(ws_url, headers).await + } +} + +impl HttpClient for OAuthSession where S: ClientAuthStore + Send + Sync + 'static, T: OAuthResolver + DpopExt + Send + Sync + 'static, + W: Send + Sync, { type Error = T::Error; @@ -417,10 +495,11 @@ where } } -impl XrpcClient for OAuthSession +impl XrpcClient for OAuthSession where S: ClientAuthStore + Send + Sync + 'static, T: OAuthResolver + DpopExt + XrpcExt + Send + Sync + 'static, + W: Send + Sync, { fn base_uri(&self) -> Url { // base_uri is a synchronous trait method; we must avoid async `.read().await`. @@ -502,10 +581,11 @@ fn is_invalid_token_response(response: &XrpcResult>) -> } } -impl IdentityResolver for OAuthSession +impl IdentityResolver for OAuthSession where S: ClientAuthStore + Send + Sync + 'static, T: OAuthResolver + IdentityResolver + XrpcExt + Send + Sync + 'static, + W: Send + Sync, { fn options(&self) -> &ResolverOptions { self.client.options() diff --git a/crates/jacquard/src/client.rs b/crates/jacquard/src/client.rs index 6b144074..f3077a4e 100644 --- a/crates/jacquard/src/client.rs +++ b/crates/jacquard/src/client.rs @@ -49,7 +49,6 @@ use jacquard_oauth::authstore::ClientAuthStore; use jacquard_oauth::client::OAuthSession; use jacquard_oauth::dpop::DpopExt; use jacquard_oauth::resolver::OAuthResolver; -use std::marker::PhantomData; use serde::Serialize; pub use token::FileAuthStore; @@ -210,10 +209,11 @@ pub trait AgentSession: XrpcClient + HttpClient + Send + Sync { fn refresh(&self) -> impl Future, ClientError>>; } -impl AgentSession for CredentialSession +impl AgentSession for CredentialSession where S: SessionStore + Send + Sync + 'static, T: IdentityResolver + HttpClient + XrpcExt + Send + Sync + 'static, + W: Send + Sync, { fn session_kind(&self) -> AgentKind { AgentKind::AppPassword @@ -227,30 +227,31 @@ where )>, > { async move { - CredentialSession::::session_info(self) + CredentialSession::::session_info(self) .await .map(|(did, sid)| (did, Some(sid))) } } fn endpoint(&self) -> impl Future { - async move { CredentialSession::::endpoint(self).await } + async move { CredentialSession::::endpoint(self).await } } fn set_options<'a>(&'a self, opts: CallOptions<'a>) -> impl Future { - async move { CredentialSession::::set_options(self, opts).await } + async move { CredentialSession::::set_options(self, opts).await } } fn refresh(&self) -> impl Future, ClientError>> { async move { - Ok(CredentialSession::::refresh(self) + Ok(CredentialSession::::refresh(self) .await? .into_static()) } } } -impl AgentSession for OAuthSession +impl AgentSession for OAuthSession where S: ClientAuthStore + Send + Sync + 'static, T: OAuthResolver + DpopExt + XrpcExt + Send + Sync + 'static, + W: Send + Sync, { fn session_kind(&self) -> AgentKind { AgentKind::OAuth @@ -264,7 +265,7 @@ where )>, > { async { - let (did, sid) = OAuthSession::::session_info(self).await; + let (did, sid) = OAuthSession::::session_info(self).await; Some((did.into_static(), Some(sid.into_static()))) } } diff --git a/crates/jacquard/src/client/credential_session.rs b/crates/jacquard/src/client/credential_session.rs index b1bc3681..f7bdce9c 100644 --- a/crates/jacquard/src/client/credential_session.rs +++ b/crates/jacquard/src/client/credential_session.rs @@ -22,6 +22,11 @@ use jacquard_identity::resolver::{ }; use std::any::Any; +#[cfg(feature = "websocket")] +use jacquard_common::websocket::{WebSocketClient, WebSocketConnection}; +#[cfg(feature = "websocket")] +use jacquard_common::xrpc::XrpcSubscription; + /// Storage key for app‑password sessions: `(account DID, session id)`. pub type SessionKey = (Did<'static>, CowStr<'static>); @@ -30,12 +35,14 @@ pub type SessionKey = (Did<'static>, CowStr<'static>); /// - Persists sessions via a pluggable `SessionStore`. /// - Automatically refreshes on token expiry. /// - Tracks a base endpoint, defaulting to the public appview until login/restore. -pub struct CredentialSession +/// - Optional WebSocket client for subscription support. +pub struct CredentialSession where S: SessionStore, { store: Arc, client: Arc, + ws_client: W, /// Default call options applied to each request (auth/headers/labelers). pub options: RwLock>, /// Active session key, if any. @@ -44,15 +51,16 @@ where pub endpoint: RwLock>, } -impl CredentialSession +impl CredentialSession where S: SessionStore, { - /// Create a new credential session using the given store and client. + /// Create a new credential session using the given store and client (no WebSocket support). pub fn new(store: Arc, client: Arc) -> Self { Self { store, client, + ws_client: (), options: RwLock::new(CallOptions::default()), key: RwLock::new(None), endpoint: RwLock::new(None), @@ -60,15 +68,33 @@ where } } -impl CredentialSession +impl CredentialSession where S: SessionStore, { + /// Create a new credential session with WebSocket client support. + pub fn new_with_ws(store: Arc, client: Arc, ws_client: W) -> Self { + Self { + store, + client, + ws_client, + options: RwLock::new(CallOptions::default()), + key: RwLock::new(None), + endpoint: RwLock::new(None), + } + } + + /// Get a reference to the WebSocket client. + pub fn ws_client(&self) -> &W { + &self.ws_client + } + /// Return a copy configured with the provided default call options. pub fn with_options(self, options: CallOptions<'_>) -> Self { Self { client: self.client, store: self.store, + ws_client: self.ws_client, options: RwLock::new(options.into_static()), key: self.key, endpoint: self.endpoint, @@ -112,7 +138,7 @@ where } } -impl CredentialSession +impl CredentialSession where S: SessionStore, T: HttpClient, @@ -150,7 +176,7 @@ where } } -impl CredentialSession +impl CredentialSession where S: SessionStore, T: HttpClient + IdentityResolver + XrpcExt + Sync + Send, @@ -385,10 +411,60 @@ where } } -impl HttpClient for CredentialSession +#[cfg(feature = "websocket")] +impl CredentialSession +where + S: SessionStore, + W: WebSocketClient, +{ + /// Subscribe to an XRPC WebSocket subscription. + /// + /// Connects to the WebSocket endpoint and threads through authentication headers. + pub async fn subscribe( + &self, + params: &Sub, + ) -> Result + where + Sub: XrpcSubscription, + { + let base_uri = self.endpoint().await; + + // Build WebSocket URL + let mut ws_url = base_uri.clone(); + ws_url.set_scheme("wss").ok(); + ws_url.set_path(&format!("/xrpc/{}", Sub::NSID)); + + // Add query params + let query_params = params.query_params(); + if !query_params.is_empty() { + let query_string = serde_html_form::to_string(&query_params) + .unwrap_or_default(); + ws_url.set_query(Some(&query_string)); + } + + // Thread auth headers (even though tokio-tungstenite-wasm doesn't support them yet) + let headers = if let Some(token) = self.access_token().await { + let auth_value = match token { + AuthorizationToken::Bearer(t) => format!("Bearer {}", t.as_ref()), + AuthorizationToken::Dpop(t) => format!("DPoP {}", t.as_ref()), + }; + vec![( + CowStr::from("Authorization"), + CowStr::from(auth_value), + )] + } else { + vec![] + }; + + self.ws_client.connect_with_headers(ws_url, headers).await + } +} + +impl HttpClient for CredentialSession where S: SessionStore + Send + Sync + 'static, T: HttpClient + XrpcExt + Send + Sync + 'static, + W: Send + Sync, { type Error = T::Error; @@ -400,10 +476,11 @@ where } } -impl XrpcClient for CredentialSession +impl XrpcClient for CredentialSession where S: SessionStore + Send + Sync + 'static, T: HttpClient + XrpcExt + Send + Sync + 'static, + W: Send + Sync, { fn base_uri(&self) -> Url { // base_uri is a synchronous trait method; avoid `.await` here. @@ -484,10 +561,11 @@ fn is_expired(response: &XrpcResult>) -> bool { } } -impl IdentityResolver for CredentialSession +impl IdentityResolver for CredentialSession where S: SessionStore + Send + Sync + 'static, T: HttpClient + IdentityResolver + Send + Sync + 'static, + W: Send + Sync, { fn options(&self) -> &ResolverOptions { self.client.options()