diff --git a/src/atproto/__init__.py b/src/atproto/__init__.py index 85fa2ba..3b4e6c9 100644 --- a/src/atproto/__init__.py +++ b/src/atproto/__init__.py @@ -1,6 +1,7 @@ +import asyncio from os import getenv from re import match as regex_match -from typing import Any +from typing import Any, TypeGuard from aiodns import DNSResolver from aiodns import error as dns_error @@ -9,6 +10,7 @@ from aiohttp.client import ClientResponse, ClientSession from src.security import is_safe_url from .kv import KV, nokv +from .types import DID, AuthserverUrl, Handle, PdsUrl from .validator import is_valid_authserver_meta PLC_DIRECTORY = getenv("PLC_DIRECTORY_URL") or "https://plc.directory" @@ -16,28 +18,46 @@ HANDLE_REGEX = r"^([a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?\.)+[a-zA-Z]([a-zA DID_REGEX = r"^did:[a-z]+:[a-zA-Z0-9._:%-]*[a-zA-Z0-9._-]$" -type AuthserverUrl = str -type PdsUrl = str -type DID = str - - -def is_valid_handle(handle: str) -> bool: +def is_valid_handle(handle: str) -> TypeGuard[Handle]: return regex_match(HANDLE_REGEX, handle) is not None -def is_valid_did(did: str) -> bool: +def is_valid_did(did: str) -> TypeGuard[DID]: return regex_match(DID_REGEX, did) is not None async def resolve_identity( client: ClientSession, query: str, - didkv: KV = nokv, -) -> tuple[str, str, dict[str, Any]] | None: + didkv: KV[Handle, DID] = nokv, + pdskv: KV[DID, PdsUrl] = nokv, +) -> tuple[DID, Handle, PdsUrl] | None: + ((done,), _) = await asyncio.wait( + ( + asyncio.create_task( + resolve_identity_microcosm(client, query, didkv=didkv, pdskv=pdskv), + name="microcosm", + ), + asyncio.create_task( + resolve_identity_raw(client, query, didkv=didkv, pdskv=pdskv), + name="raw", + ), + ), + return_when=asyncio.FIRST_COMPLETED, + ) + return done.result() + + +async def resolve_identity_raw( + client: ClientSession, + query: str, + didkv: KV[Handle, DID], + pdskv: KV[DID, PdsUrl], +) -> tuple[DID, Handle, PdsUrl] | None: """Resolves an identity to a DID, handle and DID document, verifies handles bi directionally.""" if is_valid_handle(query): - handle = query.lower() + handle = Handle(query.lower()) did = await resolve_did_from_handle(client, handle, didkv) if not did: return None @@ -47,9 +67,8 @@ async def resolve_identity( doc_handle = handle_from_doc(doc) if not doc_handle or doc_handle != handle: return None - return (did, handle, doc) - if is_valid_did(query): + elif is_valid_did(query): did = query doc = await resolve_doc_from_did(client, did) if not doc: @@ -59,12 +78,40 @@ async def resolve_identity( return None if await resolve_did_from_handle(client, handle, didkv) != did: return None - return (did, handle, doc) - return None + else: + return None + + pds_url = pds_endpoint_from_doc(doc) + if not pds_url: + return None + pdskv.set(did, value=pds_url) + + return (did, handle, pds_url) -def handle_from_doc(doc: dict[str, list[str]]) -> str | None: +async def resolve_identity_microcosm( + client: ClientSession, + query: str, + didkv: KV[Handle, DID], + pdskv: KV[DID, PdsUrl], +) -> tuple[DID, Handle, PdsUrl] | None: + base = "https://slingshot.microcosm.blue" + url = f"{base}/xrpc/com.bad-example.identity.resolveMiniDoc?identifier={query}" + response = await client.get(url) + if not response.ok: + return None + mini_doc: dict[str, str] = await response.json() + did, handle, pds = mini_doc["did"], mini_doc["handle"], mini_doc["pds"] + assert is_valid_did(did) + assert is_valid_handle(handle) + didkv.set(handle, value=did) + pds = PdsUrl(pds) + pdskv.set(did, value=pds) + return did, handle, pds + + +def handle_from_doc(doc: dict[str, list[str]]) -> Handle | None: """Return all possible handles inside the DID document.""" for aka in doc.get("alsoKnownAs", []): @@ -78,9 +125,9 @@ def handle_from_doc(doc: dict[str, list[str]]) -> str | None: async def resolve_did_from_handle( client: ClientSession, handle: str, - kv: KV = nokv, + kv: KV[Handle, DID] = nokv, reload: bool = False, -) -> str | None: +) -> DID | None: """Returns the DID for a given handle.""" if not is_valid_handle(handle): @@ -96,7 +143,7 @@ async def resolve_did_from_handle( if did is not None and is_valid_did(did): kv.set(handle, value=did) - return did + return DID(did) return None @@ -130,12 +177,14 @@ async def _resolve_did_from_handle_dns(handle: str) -> str | None: return None -def pds_endpoint_from_doc(doc: dict[str, list[dict[str, str]]]) -> str | None: +def pds_endpoint_from_doc(doc: dict[str, list[dict[str, str]]]) -> PdsUrl | None: """Returns the PDS endpoint from the DID document.""" for service in doc.get("service", []): if service.get("id") == "#atproto_pds": - return service.get("serviceEndpoint") + url = service.get("serviceEndpoint") + if url is not None: + return PdsUrl(url) return None @@ -199,7 +248,7 @@ async def resolve_authserver_from_pds( parsed: dict[str, list[str]] = await response.json() authserver_url = parsed["authorization_servers"][0] kv.set(pds_url, value=authserver_url) - return authserver_url + return AuthserverUrl(authserver_url) async def fetch_authserver_meta( diff --git a/src/atproto/kv.py b/src/atproto/kv.py index 5a0a2e1..951b080 100644 --- a/src/atproto/kv.py +++ b/src/atproto/kv.py @@ -1,15 +1,18 @@ from abc import ABC, abstractmethod from logging import Logger -from typing import override +from typing import Generic, TypeVar, override +K = TypeVar("K", bound=str) +V = TypeVar("V", bound=str) -class KV(ABC): + +class KV(ABC, Generic[K, V]): @abstractmethod - def get(self, key: str) -> str | None: + def get(self, key: K) -> V | None: pass @abstractmethod - def set(self, key: str, value: str): + def set(self, key: K, value: V): pass diff --git a/src/atproto/types.py b/src/atproto/types.py index 26e8dbf..a6f7873 100644 --- a/src/atproto/types.py +++ b/src/atproto/types.py @@ -1,12 +1,17 @@ -from typing import NamedTuple +from typing import NamedTuple, NewType + +AuthserverUrl = NewType("AuthserverUrl", str) +PdsUrl = NewType("PdsUrl", str) +Handle = NewType("Handle", str) +DID = NewType("DID", str) class OAuthAuthRequest(NamedTuple): state: str authserver_iss: str - did: str | None - handle: str | None - pds_url: str | None + did: DID | None + handle: Handle | None + pds_url: PdsUrl | None pkce_verifier: str scope: str dpop_authserver_nonce: str @@ -14,9 +19,9 @@ class OAuthAuthRequest(NamedTuple): class OAuthSession(NamedTuple): - did: str - handle: str | None - pds_url: str + did: DID + handle: Handle | None + pds_url: PdsUrl authserver_iss: str access_token: str | None refresh_token: str | None diff --git a/src/db.py b/src/db.py index e442a10..985939e 100644 --- a/src/db.py +++ b/src/db.py @@ -1,14 +1,15 @@ import sqlite3 from logging import Logger from sqlite3 import Connection -from typing import override +from typing import Generic, cast, override from flask import Flask, g from src.atproto.kv import KV as BaseKV +from src.atproto.kv import K, V -class KV(BaseKV): +class KV(BaseKV, Generic[K, V]): db: Connection logger: Logger prefix: str @@ -19,7 +20,7 @@ class KV(BaseKV): self.prefix = prefix @override - def get(self, key: str) -> str | None: + def get(self, key: K) -> V | None: cursor = self.db.cursor() row: dict[str, str] | None = cursor.execute( "select value from keyval where prefix = ? and key = ?", @@ -27,11 +28,11 @@ class KV(BaseKV): ).fetchone() if row is not None: self.logger.debug(f"returning cached {self.prefix}({key})") - return row["value"] + return cast(V, row["value"]) return None @override - def set(self, key: str, value: str): + def set(self, key: K, value: V): self.logger.debug(f"caching {self.prefix}({key}): {value}") cursor = self.db.cursor() _ = cursor.execute( diff --git a/src/main.py b/src/main.py index 714f5eb..35576bd 100644 --- a/src/main.py +++ b/src/main.py @@ -8,14 +8,13 @@ from flask_htmx import HTMX from flask_htmx import make_response as htmx_response from src.atproto import ( - PdsUrl, get_record, is_valid_did, resolve_did_from_handle, resolve_pds_from_did, ) from src.atproto.oauth import pds_authed_req -from src.atproto.types import OAuthSession +from src.atproto.types import DID, Handle, OAuthSession, PdsUrl from src.auth import get_auth_session, save_auth_session from src.db import KV, close_db_connection, get_db, init_db from src.oauth import oauth @@ -52,8 +51,8 @@ async def page_profile(atid: str): reload = request.args.get("reload") is not None db = get_db(app) - didkv = KV(db, app.logger, "did_from_handle") - pdskv = KV(db, app.logger, "pds_from_did") + didkv = KV[Handle, DID](db, app.logger, "did_from_handle") + pdskv = KV[DID, PdsUrl](db, app.logger, "pds_from_did") async with ClientSession() as client: if atid.startswith("@"): diff --git a/src/oauth.py b/src/oauth.py index abc7b71..f7b0cc9 100644 --- a/src/oauth.py +++ b/src/oauth.py @@ -9,12 +9,18 @@ from src.atproto import ( fetch_authserver_meta, is_valid_did, is_valid_handle, - pds_endpoint_from_doc, resolve_authserver_from_pds, resolve_identity, ) from src.atproto.oauth import initial_token_request, send_par_auth_request -from src.atproto.types import OAuthAuthRequest, OAuthSession +from src.atproto.types import ( + DID, + AuthserverUrl, + Handle, + OAuthAuthRequest, + OAuthSession, + PdsUrl, +) from src.auth import ( delete_auth_request, get_auth_request, @@ -35,22 +41,26 @@ async def oauth_start(): return redirect(url_for("page_login"), 303) db = get_db(current_app) - pdskv = KV(db, current_app.logger, "authserver_from_pds") + didkv = KV[Handle, DID](db, current_app.logger, "did_from_handle") + pdskv = KV[DID, PdsUrl](db, current_app.logger, "pds_from_did") + authserverkv = KV[PdsUrl, AuthserverUrl]( + db, + current_app.logger, + "authserver_from_pds", + ) client = ClientSession() if is_valid_handle(username) or is_valid_did(username): login_hint = username - kv = KV(db, current_app.logger, "did_from_handle") - identity = await resolve_identity(client, username, didkv=kv) + identity = await resolve_identity(client, username, didkv=didkv, pdskv=pdskv) if identity is None: return "couldnt resolve identity", 500 - did, handle, doc = identity - pds_url = pds_endpoint_from_doc(doc) - if not pds_url: - return "pds not found", 404 + did, handle, pds_url = identity current_app.logger.debug(f"account PDS: {pds_url}") - authserver_url = await resolve_authserver_from_pds(client, pds_url, pdskv) + authserver_url = await resolve_authserver_from_pds( + client, pds_url, authserverkv + ) if not authserver_url: return "authserver not found", 404 @@ -58,7 +68,8 @@ async def oauth_start(): did, handle, pds_url = None, None, None login_hint = None authserver_url = ( - await resolve_authserver_from_pds(client, username, pdskv) or username + await resolve_authserver_from_pds(client, PdsUrl(username), authserverkv) + or username ) else: @@ -175,10 +186,7 @@ async def oauth_callback(): identity = await resolve_identity(client, did, didkv=didkv) if not identity: return "could not resolve identity", 500 - did, handle, did_doc = identity - pds_url = pds_endpoint_from_doc(did_doc) - if not pds_url: - return "could not resolve pds", 500 + did, handle, pds_url = identity authserver_url = await resolve_authserver_from_pds( client, pds_url,