from dataclasses import dataclass from typing import Any, TypeVar import httpx from atproto.store import AtprotoStore, IdentityInfo, Session from cross.media import Blob from util.util import LOGGER, normalize_service_url class XRPCError(Exception): def __init__( self, message: str, status_code: int | None = None, response_data: dict[str, Any] | None = None, ) -> None: super().__init__(message) self.status_code = status_code self.response_data = response_data T = TypeVar("T") @dataclass class XRPCResponse: data: dict[str, Any] status_code: int headers: dict[str, str] class XRPCClient: def __init__( self, pds_url: str, store: AtprotoStore, http: httpx.Client | None = None, identifier: str | None = None, password: str | None = None, ) -> None: self.pds_url: str = normalize_service_url(pds_url) self.store: AtprotoStore = store self.http: httpx.Client = http if http else httpx.Client() if identifier and password: self.login(identifier, password) def login(self, identifier: str, password: str) -> Session: cached = self.store.get_session_by_pds(self.pds_url, identifier) if cached and not cached.is_refresh_token_expired(): return cached return self.create_session(identifier, password) def get_session(self, did: str) -> Session | None: session = self.store.get_session(did) if not session: return None if session.is_access_token_expired(): if not session.is_refresh_token_expired(): LOGGER.info("Refreshing session for '%s'", session.did) return self.refresh_session(session) self.store.remove_session(did) raise ValueError( "Both access and refresh tokens expired. Please login again." ) return session def create_session( self, identifier: str, password: str, auth_factor_token: str | None = None, ) -> Session: url = f"{self.pds_url}/xrpc/com.atproto.server.createSession" payload: dict[str, str] = {"identifier": identifier, "password": password} if auth_factor_token: payload["authFactorToken"] = auth_factor_token response = self.http.post(url, json=payload, timeout=30) match response.status_code: case 200: pass case 401: raise ValueError("Invalid identifier or password") case 400: raise ValueError(f"Authentication failed: {response.json()}") case _: raise ValueError( f"Authentication failed with status {response.status_code}" ) session = Session.from_dict(response.json(), self.pds_url) self.store.set_session(session) LOGGER.info("Created session for '%s'", session.did) return session def refresh_session(self, session: Session) -> Session: url = f"{self.pds_url}/xrpc/com.atproto.server.refreshSession" headers = {"Authorization": f"Bearer {session.refresh_jwt}"} response = self.http.post(url, headers=headers, timeout=30) match response.status_code: case 200: pass case 401: error_data = response.json() if response.content else {} raise ValueError(f"Refresh failed: {error_data}") case 400: raise ValueError(f"Refresh failed: {response.json()}") case _: raise ValueError(f"Refresh failed with status {response.status_code}") new_session = Session.from_dict(response.json(), self.pds_url) self.store.set_session(new_session) LOGGER.info("Refreshed session for '%s'", new_session.did) return new_session def get_access_token(self, did: str) -> str | None: session = self.get_session(did) return session.access_jwt if session else None def call( self, method: str, params: dict[str, Any] | None = None, data: dict[str, Any] | None = None, did: str | None = None, ) -> XRPCResponse: url = f"{self.pds_url}/xrpc/{method}" headers = ( { "Authorization": f"Bearer {self.get_access_token(did)}", "Content-Type": "application/json", } if did else {} ) if params and data: raise ValueError("Cannot specify both params and data") try: if params: response = self.http.get( url, params=params, headers=headers, timeout=30 ) elif data: response = self.http.post(url, json=data, headers=headers, timeout=30) else: response = self.http.get(url, headers=headers, timeout=30) except httpx.RequestError as e: raise XRPCError(f"Request failed: {e}") from e try: response_data = response.json() if response.content else {} except ValueError as e: raise XRPCError( f"Invalid JSON response: {e}", status_code=response.status_code ) from e if response.status_code >= 400: error_msg = response_data.get( "message", f"Request failed with status {response.status_code}" ) raise XRPCError( error_msg, status_code=response.status_code, response_data=response_data ) return XRPCResponse( data=response_data, status_code=response.status_code, headers=dict(response.headers), ) def upload_blob( self, blob: Blob, content_type: str, did: str, ) -> dict[str, Any]: token = self.get_access_token(did) if not token: raise ValueError(f"No valid session found for {did}") url = f"{self.pds_url}/xrpc/com.atproto.repo.uploadBlob" headers = { "Authorization": f"Bearer {token}", "Content-Type": content_type, } try: with open(blob.path, "rb") as f: response = self.http.post(url, content=f, headers=headers, timeout=60) except httpx.RequestError as e: raise XRPCError(f"Blob upload request failed: {e}") from e except OSError as e: raise XRPCError(f"Could not read blob file: {e}") from e if response.status_code != 200: error_data = response.json() if response.content else {} raise XRPCError( f"Blob upload failed: {response.status_code}", status_code=response.status_code, response_data=error_data, ) try: result: dict[str, Any] = response.json() except ValueError as e: raise XRPCError(f"Invalid JSON response from blob upload: {e}") from e return result def resolve_identity( identifier: str, store: AtprotoStore, http: httpx.Client | None = None ) -> IdentityInfo: import env cached = store.get_identity(identifier) if cached: return cached url = f"{env.SLINGSHOT_URL}/xrpc/com.bad-example.identity.resolveMiniDoc" client = http if http else httpx.Client() try: response = client.get(url, params={"identifier": identifier}, timeout=10) response.raise_for_status() try: data = response.json() except ValueError as e: raise ValueError( f"Invalid JSON response from identity resolver: {e}" ) from e match response.status_code: case 200: identity = IdentityInfo.from_dict(data) store.set_identity(identifier, identity) return identity case 404: raise ValueError(f"Identity not found: {identifier}") case _: error_msg = data.get( "message", f"Identity resolver returned status {response.status_code}", ) raise ValueError(error_msg) except httpx.RequestError as e: raise ValueError(f"Failed to resolve identity {identifier}: {e}") from e