diff --git a/crates/jacquard-common/src/session.rs b/crates/jacquard-common/src/session.rs index 0d0f7430..a788ec09 100644 --- a/crates/jacquard-common/src/session.rs +++ b/crates/jacquard-common/src/session.rs @@ -25,7 +25,7 @@ pub trait Session { async fn access_token(&self) -> Result; - async fn refresh(&self) -> Result<(), SessionStoreError>; + async fn refresh(&self) -> Result; } /// Errors emitted by session stores. diff --git a/crates/jacquard-identity/src/lib.rs b/crates/jacquard-identity/src/lib.rs index 9186a550..9c41f802 100644 --- a/crates/jacquard-identity/src/lib.rs +++ b/crates/jacquard-identity/src/lib.rs @@ -36,14 +36,14 @@ use url::{ParseError, Url}; use hickory_resolver::{TokioAsyncResolver, config::ResolverConfig}; /// Default resolver implementation with configurable fallback order. -pub struct DefaultResolver { +pub struct JacquardResolver { http: reqwest::Client, opts: ResolverOptions, #[cfg(feature = "dns")] dns: Option, } -impl DefaultResolver { +impl JacquardResolver { /// Create a new instance of the default resolver with all options (except DNS) up front pub fn new(http: reqwest::Client, opts: ResolverOptions) -> Self { Self { @@ -189,7 +189,7 @@ impl DefaultResolver { } } -impl DefaultResolver { +impl JacquardResolver { /// Resolve handle to DID via a PDS XRPC call (stateless, unauth by default) pub async fn resolve_handle_via_pds( &self, @@ -270,7 +270,7 @@ impl DefaultResolver { } #[async_trait::async_trait] -impl IdentityResolver for DefaultResolver { +impl IdentityResolver for JacquardResolver { fn options(&self) -> &ResolverOptions { &self.opts } @@ -417,7 +417,7 @@ impl IdentityResolver for DefaultResolver { } } -impl HttpClient for DefaultResolver { +impl HttpClient for JacquardResolver { async fn send_http( &self, request: http::Request>, @@ -438,7 +438,7 @@ pub enum IdentityWarning { }, } -impl DefaultResolver { +impl JacquardResolver { /// Resolve a handle to its DID, fetch the DID document, and return doc plus any warnings. /// This applies the default equality check on the document id (error with doc if mismatch). pub async fn resolve_handle_and_doc( @@ -523,7 +523,7 @@ impl MiniDocResponse { } /// Resolver specialized for unauthenticated/public flows using reqwest and stateless XRPC -pub type PublicResolver = DefaultResolver; +pub type PublicResolver = JacquardResolver; impl Default for PublicResolver { /// Build a resolver with: @@ -539,7 +539,7 @@ impl Default for PublicResolver { fn default() -> Self { let http = reqwest::Client::new(); let opts = ResolverOptions::default(); - let resolver = DefaultResolver::new(http, opts); + let resolver = JacquardResolver::new(http, opts); #[cfg(feature = "dns")] let resolver = resolver.with_system_dns(); resolver @@ -552,7 +552,7 @@ pub fn slingshot_resolver_default() -> PublicResolver { let http = reqwest::Client::new(); let mut opts = ResolverOptions::default(); opts.plc_source = PlcSource::slingshot_default(); - let resolver = DefaultResolver::new(http, opts); + let resolver = JacquardResolver::new(http, opts); #[cfg(feature = "dns")] let resolver = resolver.with_system_dns(); resolver @@ -564,7 +564,7 @@ mod tests { #[test] fn did_web_urls() { - let r = DefaultResolver::new(reqwest::Client::new(), ResolverOptions::default()); + let r = JacquardResolver::new(reqwest::Client::new(), ResolverOptions::default()); assert_eq!( r.test_did_web_url_raw("did:web:example.com"), "https://example.com/.well-known/did.json" @@ -577,7 +577,7 @@ mod tests { #[test] fn slingshot_mini_doc_url_build() { - let r = DefaultResolver::new(reqwest::Client::new(), ResolverOptions::default()); + let r = JacquardResolver::new(reqwest::Client::new(), ResolverOptions::default()); let base = Url::parse("https://slingshot.microcosm.blue").unwrap(); let url = r.slingshot_mini_doc_url(&base, "bad-example.com").unwrap(); assert_eq!( diff --git a/crates/jacquard-oauth/src/client.rs b/crates/jacquard-oauth/src/client.rs index ea253fa1..540057fb 100644 --- a/crates/jacquard-oauth/src/client.rs +++ b/crates/jacquard-oauth/src/client.rs @@ -1,3 +1,14 @@ +use crate::{ + atproto::atproto_client_metadata, + authstore::ClientAuthStore, + dpop::DpopExt, + error::{OAuthError, Result}, + request::{OAuthMetadata, exchange_code, par}, + resolver::OAuthResolver, + scopes::Scope, + session::{ClientData, ClientSessionData, DpopClientData, SessionRegistry}, + types::{AuthorizeOptions, CallbackParams}, +}; use jacquard_common::{ AuthorizationToken, CowStr, IntoStatic, error::{AuthError, ClientError, TransportError, XrpcResult}, @@ -8,23 +19,10 @@ use jacquard_common::{ }, }; use jose_jwk::JwkSet; -use smol_str::SmolStr; use std::sync::Arc; use tokio::sync::RwLock; use url::Url; -use crate::{ - atproto::atproto_client_metadata, - authstore::ClientAuthStore, - dpop::DpopExt, - error::{OAuthError, Result}, - request::{OAuthMetadata, exchange_code, par}, - resolver::OAuthResolver, - scopes::Scope, - session::{ClientData, ClientSessionData, DpopClientData, SessionRegistry}, - types::{AuthorizeOptions, CallbackParams}, -}; - pub struct OAuthClient where T: OAuthResolver, @@ -242,7 +240,7 @@ where (data.account_did.clone(), data.session_id.clone()) } - pub async fn pds(&self) -> Url { + pub async fn endpoint(&self) -> Url { self.data.read().await.host_url.clone() } diff --git a/crates/jacquard/src/client.rs b/crates/jacquard/src/client.rs index 9309389d..8877ecd3 100644 --- a/crates/jacquard/src/client.rs +++ b/crates/jacquard/src/client.rs @@ -4,7 +4,7 @@ //! client implementation that manages session tokens. mod at_client; - +pub mod credential_session; mod token; pub use at_client::{AtClient, SendOverrides}; @@ -21,8 +21,6 @@ use jacquard_common::{ pub use token::FileAuthStore; use url::Url; -// Note: Stateless and stateful XRPC clients are implemented in xrpc_call.rs and at_client.rs - pub(crate) const NSID_REFRESH_SESSION: &str = "com.atproto.server.refreshSession"; /// Basic client wrapper: reqwest transport + in-memory session store. diff --git a/crates/jacquard/src/client/credential_session.rs b/crates/jacquard/src/client/credential_session.rs new file mode 100644 index 00000000..e9347e20 --- /dev/null +++ b/crates/jacquard/src/client/credential_session.rs @@ -0,0 +1,189 @@ +use std::sync::Arc; + +use jacquard_api::com_atproto::server::refresh_session::RefreshSession; +use jacquard_common::{ + AuthorizationToken, CowStr, IntoStatic, + error::{AuthError, ClientError, XrpcResult}, + http_client::HttpClient, + session::SessionStore, + types::{ + did::Did, + xrpc::{CallOptions, Response, XrpcClient, XrpcError, XrpcExt, XrpcRequest}, + }, +}; +use tokio::sync::RwLock; +use url::Url; + +use crate::client::{AtpSession, token::StoredSession}; + +pub type SessionKey = (Did<'static>, CowStr<'static>); + +pub struct CredentialSession +where + S: SessionStore, +{ + store: Arc, + client: Arc, + pub options: RwLock>, + pub key: RwLock>, + pub endpoint: RwLock>, +} + +impl CredentialSession +where + S: SessionStore, +{ + pub fn new(store: Arc, client: Arc) -> Self { + Self { + store, + client, + options: RwLock::new(CallOptions::default()), + key: RwLock::new(None), + endpoint: RwLock::new(None), + } + } +} + +impl CredentialSession +where + S: SessionStore, +{ + pub fn with_options(self, options: CallOptions<'_>) -> Self { + Self { + client: self.client, + store: self.store, + options: RwLock::new(options.into_static()), + key: self.key, + endpoint: self.endpoint, + } + } + + pub async fn set_options(&self, options: CallOptions<'_>) { + *self.options.write().await = options.into_static(); + } + + pub async fn session_info(&self) -> Option { + self.key.read().await.clone() + } + + pub async fn endpoint(&self) -> Url { + self.endpoint.read().await.clone().unwrap_or( + Url::parse("https://public.bsky.app").expect("public appview should be valid url"), + ) + } + + pub async fn set_endpoint(&self, endpoint: Url) { + *self.endpoint.write().await = Some(endpoint); + } + + pub async fn access_token(&self) -> Option> { + let key = self.key.read().await.clone()?; + let session = self.store.get(&key).await; + session.map(|session| AuthorizationToken::Bearer(session.access_jwt)) + } + + pub async fn refresh_token(&self) -> Option> { + let key = self.key.read().await.clone()?; + let session = self.store.get(&key).await; + session.map(|session| AuthorizationToken::Bearer(session.refresh_jwt)) + } +} + +impl CredentialSession +where + S: SessionStore, + T: HttpClient, +{ + pub async fn refresh(&self) -> Result, ClientError> { + let key = self.key.read().await.clone().ok_or(ClientError::Auth( + jacquard_common::error::AuthError::NotAuthenticated, + ))?; + let session = self.store.get(&key).await; + let endpoint = self.endpoint().await; + let mut opts = self.options.read().await.clone(); + opts.auth = session.map(|s| AuthorizationToken::Bearer(s.refresh_jwt)); + let response = self + .client + .xrpc(endpoint) + .with_options(opts) + .send(&RefreshSession) + .await?; + let refresh = response + .into_output() + .map_err(|_| ClientError::Auth(jacquard_common::error::AuthError::RefreshFailed))?; + + let new_session: AtpSession = refresh.into(); + let token = AuthorizationToken::Bearer(new_session.access_jwt.clone()); + self.store + .set(key, new_session) + .await + .map_err(|_| ClientError::Auth(jacquard_common::error::AuthError::RefreshFailed))?; + + Ok(token) + } +} + +impl HttpClient for CredentialSession +where + S: SessionStore + Send + Sync + 'static, + T: HttpClient + XrpcExt + Send + Sync + 'static, +{ + type Error = T::Error; + + async fn send_http( + &self, + request: http::Request>, + ) -> core::result::Result>, Self::Error> { + self.client.send_http(request).await + } +} + +impl XrpcClient for CredentialSession +where + S: SessionStore + Send + Sync + 'static, + T: HttpClient + XrpcExt + Send + Sync + 'static, +{ + fn base_uri(&self) -> Url { + self.endpoint.blocking_read().clone().unwrap_or( + Url::parse("https://public.bsky.app").expect("public appview should be valid url"), + ) + } + async fn send( + self, + request: &R, + ) -> XrpcResult> { + let base_uri = self.base_uri(); + let auth = self.access_token().await; + let mut opts = self.options.read().await.clone(); + opts.auth = auth; + let resp = self + .client + .xrpc(base_uri.clone()) + .with_options(opts.clone()) + .send(request) + .await; + + if is_expired(&resp) { + let auth = self.refresh().await?; + opts.auth = Some(auth); + self.client + .xrpc(base_uri) + .with_options(opts) + .send(request) + .await + } else { + resp + } + } +} + +fn is_expired(response: &XrpcResult>) -> bool { + match response { + Err(ClientError::Auth(AuthError::TokenExpired)) => true, + Ok(resp) => match resp.parse() { + Err(XrpcError::Auth(AuthError::TokenExpired)) => true, + _ => false, + }, + _ => false, + } +} diff --git a/crates/jacquard/src/client/token.rs b/crates/jacquard/src/client/token.rs index f2b74927..4d863c46 100644 --- a/crates/jacquard/src/client/token.rs +++ b/crates/jacquard/src/client/token.rs @@ -15,7 +15,7 @@ use std::path::{Path, PathBuf}; use url::Url; #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] -enum StoredSession { +pub enum StoredSession { Atp(StoredAtSession), OAuth(OAuthSession), OAuthState(OAuthState),