diff --git a/crates/jacquard-common/src/types/language.rs b/crates/jacquard-common/src/types/language.rs index e5d796b8..b9c5c8fd 100644 --- a/crates/jacquard-common/src/types/language.rs +++ b/crates/jacquard-common/src/types/language.rs @@ -1,3 +1,4 @@ +use langtag::InvalidLangTag; use serde::{Deserialize, Deserializer, Serialize, de::Error}; use smol_str::{SmolStr, ToSmolStr}; use std::fmt; @@ -36,6 +37,11 @@ impl Language { Ok(Language(SmolStr::new_static(tag.as_str()))) } + fn new_owned(lang: SmolStr) -> Result { + let tag = langtag::LangTag::new(&lang).map_err(|e| e.to_smolstr())?; + Ok(Language(SmolStr::new(tag.as_str()))) + } + /// Infallible constructor for when you *know* the string is a valid IETF language tag. /// Will panic on invalid tag. If you're manually decoding atproto records /// or API values you know are valid (rather than using serde), this is the one to use. @@ -75,8 +81,8 @@ impl<'de> Deserialize<'de> for Language { where D: Deserializer<'de>, { - let value: &str = Deserialize::deserialize(deserializer)?; - Self::new(value).map_err(D::Error::custom) + let value = Deserialize::deserialize(deserializer)?; + Self::new_owned(value).map_err(D::Error::custom) } } diff --git a/crates/jacquard/src/client.rs b/crates/jacquard/src/client.rs index 619923d7..af5aa47b 100644 --- a/crates/jacquard/src/client.rs +++ b/crates/jacquard/src/client.rs @@ -452,10 +452,11 @@ impl MemoryCredentialSession { identifier: CowStr<'_>, password: CowStr<'_>, session_id: Option>, + pds: Option, ) -> ClientResult<(Self, AtpSession)> { let session = MemoryCredentialSession::unauthenticated(); let auth = session - .login(identifier, password, session_id, None, None) + .login(identifier, password, session_id, None, None, pds) .await?; Ok((session, auth)) } diff --git a/crates/jacquard/src/client/credential_session.rs b/crates/jacquard/src/client/credential_session.rs index d2c4d30e..54c2bcce 100644 --- a/crates/jacquard/src/client/credential_session.rs +++ b/crates/jacquard/src/client/credential_session.rs @@ -204,6 +204,7 @@ where session_id: Option>, allow_takendown: Option, auth_factor_token: Option>, + pds: Option, ) -> std::result::Result where S: Any + 'static, @@ -213,7 +214,9 @@ where tracing::info_span!("credential_session_login", identifier = %identifier).entered(); // Resolve PDS base - let pds = if identifier.as_ref().starts_with("http://") + let pds = if let Some(pds) = pds { + pds + } else if identifier.as_ref().starts_with("http://") || identifier.as_ref().starts_with("https://") { Url::parse(identifier.as_ref()).map_err(|e: url::ParseError| { @@ -232,6 +235,12 @@ where ClientError::invalid_request("missing PDS endpoint") .with_help("DID document must include a PDS service endpoint") })? + } else if identifier.as_ref().contains("@") && !identifier.as_ref().starts_with("@") { + // we're going to assume its an email + pds.ok_or_else(|| { + ClientError::invalid_request("missing PDS endpoint") + .with_help("When logging in with email, we need your PDS") + })? } else { // treat as handle let handle = diff --git a/crates/jacquard/tests/credential_session.rs b/crates/jacquard/tests/credential_session.rs index 9dc964ef..5470a265 100644 --- a/crates/jacquard/tests/credential_session.rs +++ b/crates/jacquard/tests/credential_session.rs @@ -179,6 +179,7 @@ async fn credential_login_and_auto_refresh() { Some(jacquard::CowStr::from("session")), None, None, + None, ) .await .expect("login ok");