diff --git a/src/atproto/__init__.py b/src/atproto/__init__.py index 086151e..e743d47 100644 --- a/src/atproto/__init__.py +++ b/src/atproto/__init__.py @@ -1,5 +1,6 @@ -import aiohttp -import aiodns +from aiodns import DNSResolver, error as dns_error +from aiohttp.client import ClientSession +from os import getenv from re import match as regex_match from typing import Any @@ -7,14 +8,14 @@ from .kv import KV, nokv from .validator import is_valid_authserver_meta from ..security import is_safe_url -PLC_DIRECTORY = "https://plc.directory" +PLC_DIRECTORY = getenv("PLC_DIRECTORY_URL") or "https://plc.directory" HANDLE_REGEX = r"^([a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?\.)+[a-zA-Z]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?$" DID_REGEX = r"^did:[a-z]+:[a-zA-Z0-9._:%-]*[a-zA-Z0-9._-]$" -AuthserverUrl = str -PdsUrl = str -DID = str +type AuthserverUrl = str +type PdsUrl = str +type DID = str def is_valid_handle(handle: str) -> bool: @@ -26,6 +27,7 @@ def is_valid_did(did: str) -> bool: async def resolve_identity( + client: ClientSession, query: str, didkv: KV = nokv, ) -> tuple[str, str, dict[str, Any]] | None: @@ -36,17 +38,17 @@ async def resolve_identity( did = await resolve_did_from_handle(handle, didkv) if not did: return None - doc = await resolve_doc_from_did(did) + doc = await resolve_doc_from_did(client, did) if not doc: return None - handles = handles_from_doc(doc) - if not handles or handle not in handles: + 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): did = query - doc = await resolve_doc_from_did(did) + doc = await resolve_doc_from_did(client, did) if not doc: return None handle = handle_from_doc(doc) @@ -59,24 +61,15 @@ async def resolve_identity( return None -def handles_from_doc(doc: dict[str, list[str]]) -> list[str]: +def handle_from_doc(doc: dict[str, list[str]]) -> str | None: """Return all possible handles inside the DID document.""" - handles: list[str] = [] + for aka in doc.get("alsoKnownAs", []): if aka.startswith("at://"): handle = aka[5:].lower() if is_valid_handle(handle): - handles.append(handle) - return handles - - -def handle_from_doc(doc: dict[str, list[str]]) -> str | None: - """Return the first handle inside the DID document.""" - handles = handles_from_doc(doc) - try: - return handles[0] - except IndexError: - return None + return handle + return None async def resolve_did_from_handle( @@ -94,10 +87,10 @@ async def resolve_did_from_handle( print(f"returning cached did for {handle}") return did - resolver = aiodns.DNSResolver() + resolver = DNSResolver() try: result = await resolver.query(f"_atproto.{handle}", "TXT") - except aiodns.error.DNSError: + except dns_error.DNSError: return None for record in result: @@ -122,6 +115,7 @@ def pds_endpoint_from_doc(doc: dict[str, list[dict[str, str]]]) -> str | None: async def resolve_pds_from_did( + client: ClientSession, did: DID, kv: KV = nokv, reload: bool = False, @@ -131,7 +125,7 @@ async def resolve_pds_from_did( print(f"returning cached pds for {did}") return pds - doc = await resolve_doc_from_did(did) + doc = await resolve_doc_from_did(client, did) if doc is None: return None pds = doc["service"][0]["serviceEndpoint"] @@ -143,24 +137,26 @@ async def resolve_pds_from_did( async def resolve_doc_from_did( + client: ClientSession, did: DID, - directory: str = PLC_DIRECTORY, ) -> dict[str, Any] | None: - async with aiohttp.ClientSession() as client: - if did.startswith("did:plc:"): - response = await client.get(f"{directory}/{did}") - if response.ok: - return await response.json() - return None + """Returns the DID document""" - if did.startswith("did:web:"): - # TODO: resolve did:web - return None + if did.startswith("did:plc:"): + response = await client.get(f"{PLC_DIRECTORY}/{did}") + if response.ok: + return await response.json() + return None + + if did.startswith("did:web:"): + # TODO: resolve did:web + raise Exception("resolve did:web") return None async def resolve_authserver_from_pds( + client: ClientSession, pds_url: PdsUrl, kv: KV = nokv, reload: bool = False, @@ -174,31 +170,34 @@ async def resolve_authserver_from_pds( assert is_safe_url(pds_url) endpoint = f"{pds_url}/.well-known/oauth-protected-resource" - async with aiohttp.ClientSession() as client: - response = await client.get(endpoint) - if response.status != 200: - return None - parsed: dict[str, list[str]] = await response.json() - authserver_url = parsed["authorization_servers"][0] - print(f"caching authserver {authserver_url} for PDS {pds_url}") - kv.set(pds_url, value=authserver_url) - return authserver_url + response = await client.get(endpoint) + if response.status != 200: + return None + parsed: dict[str, list[str]] = await response.json() + authserver_url = parsed["authorization_servers"][0] + print(f"caching authserver {authserver_url} for PDS {pds_url}") + kv.set(pds_url, value=authserver_url) + return authserver_url -async def fetch_authserver_meta(authserver_url: str) -> dict[str, str] | None: +async def fetch_authserver_meta( + client: ClientSession, + authserver_url: str, +) -> dict[str, str] | None: """Returns metadata from the authserver""" + assert is_safe_url(authserver_url) endpoint = f"{authserver_url}/.well-known/oauth-authorization-server" - async with aiohttp.ClientSession() as client: - response = await client.get(endpoint) - if not response.ok: - return None - meta: dict[str, Any] = await response.json() - assert is_valid_authserver_meta(meta, authserver_url) - return meta + response = await client.get(endpoint) + if not response.ok: + return None + meta: dict[str, Any] = await response.json() + assert is_valid_authserver_meta(meta, authserver_url) + return meta async def get_record( + client: ClientSession, pds: str, repo: str, collection: str, @@ -207,16 +206,14 @@ async def get_record( ) -> dict[str, Any] | None: """Retrieve record from PDS. Verifies type is the same as collection name.""" - async with aiohttp.ClientSession() as client: - response = await client.get( - f"{pds}/xrpc/com.atproto.repo.getRecord?repo={repo}&collection={collection}&rkey={record}" - ) - if not response.ok: - return None - parsed = await response.json() - value: dict[str, Any] = parsed["value"] - if value["$type"] != (type or collection): - return None - del value["$type"] + params = {"repo": repo, "collection": collection, "rkey": record} + response = await client.get(f"{pds}/xrpc/com.atproto.repo.getRecord", params=params) + if not response.ok: + return None + parsed = await response.json() + value: dict[str, Any] = parsed["value"] + if value["$type"] != (type or collection): + return None + del value["$type"] - return value + return value diff --git a/src/atproto/oauth.py b/src/atproto/oauth.py index 88b8e51..ee7ea0a 100644 --- a/src/atproto/oauth.py +++ b/src/atproto/oauth.py @@ -1,10 +1,10 @@ from typing import Any, Callable, NamedTuple import time import json +from aiohttp.client import ClientSession, ClientResponse from authlib.jose import JsonWebKey, Key, jwt from authlib.common.security import generate_token from authlib.oauth2.rfc7636 import create_s256_code_challenge -from aiohttp import ClientResponse from . import fetch_authserver_meta @@ -84,7 +84,6 @@ async def send_par_auth_request( respjson = await resp.json() if resp.status == 400 and respjson["error"] == "use_dpop_nonce": dpop_authserver_nonce = resp.headers["DPoP-Nonce"] - print(f"retrying with new auth server DPoP nonce: {dpop_authserver_nonce}") dpop_proof = _authserver_dpop_jwt( "POST", par_url, dpop_authserver_nonce, dpop_private_jwk ) @@ -105,6 +104,7 @@ async def send_par_auth_request( # Returns token response (OAuthTokens) and DPoP nonce (str) # IMPORTANT: the 'tokens.sub' field must be verified against the original request by code calling this function. async def initial_token_request( + client: ClientSession, auth_request: OAuthAuthRequest, code: str, app_url: str, @@ -113,7 +113,7 @@ async def initial_token_request( authserver_url = auth_request.authserver_iss # Re-fetch server metadata - authserver_meta = await fetch_authserver_meta(authserver_url) + authserver_meta = await fetch_authserver_meta(client, authserver_url) if not authserver_meta: raise Exception("missing authserver meta") @@ -146,24 +146,16 @@ async def initial_token_request( # IMPORTANT: Token URL is untrusted input, SSRF mitigations are needed assert is_safe_url(token_url) - async with hardened_http.get_session() as session: - resp = await session.post(token_url, data=params, headers={"DPoP": dpop_proof}) + resp = await client.post(token_url, data=params, headers={"DPoP": dpop_proof}) # Handle DPoP missing/invalid nonce error by retrying with server-provided nonce respjson = await resp.json() if resp.status == 400 and respjson["error"] == "use_dpop_nonce": dpop_authserver_nonce = resp.headers["DPoP-Nonce"] - print(f"retrying with new auth server DPoP nonce: {dpop_authserver_nonce}") - # print(server_nonce) dpop_proof = _authserver_dpop_jwt( "POST", token_url, dpop_authserver_nonce, dpop_private_jwk ) - async with hardened_http.get_session() as session: - resp = await session.post( - token_url, - data=params, - headers={"DPoP": dpop_proof}, - ) + resp = await client.post(token_url, data=params, headers={"DPoP": dpop_proof}) resp.raise_for_status() token_body = await resp.json() @@ -174,6 +166,7 @@ async def initial_token_request( # Returns token response (OAuthTokens) and DPoP nonce (str) async def refresh_token_request( + client: ClientSession, user: OAuthSession, app_url: str, client_secret_jwk: Key, @@ -181,7 +174,7 @@ async def refresh_token_request( authserver_url = user.authserver_iss # Re-fetch server metadata - authserver_meta = await fetch_authserver_meta(authserver_url) + authserver_meta = await fetch_authserver_meta(client, authserver_url) if not authserver_meta: raise Exception("missing authserver meta") @@ -218,8 +211,6 @@ async def refresh_token_request( respjson = await resp.json() if resp.status == 400 and respjson["error"] == "use_dpop_nonce": dpop_authserver_nonce = resp.headers["DPoP-Nonce"] - print(f"retrying with new auth server DPoP nonce: {dpop_authserver_nonce}") - # print(server_nonce) dpop_proof = _authserver_dpop_jwt( "POST", token_url, dpop_authserver_nonce, dpop_private_jwk ) @@ -278,7 +269,6 @@ async def pds_authed_req( respjson = await response.json() if response.status in [400, 401] and respjson["error"] == "use_dpop_nonce": dpop_pds_nonce = response.headers["DPoP-Nonce"] - print(f"retrying with new PDS DPoP nonce: {dpop_pds_nonce}") update_dpop_pds_nonce(dpop_pds_nonce) continue break diff --git a/src/main.py b/src/main.py index 7e81183..a1f5f0c 100644 --- a/src/main.py +++ b/src/main.py @@ -1,6 +1,7 @@ import asyncio import json +from aiohttp.client import ClientSession from flask import Flask, g, session, redirect, render_template, request, url_for from typing import Any @@ -12,7 +13,7 @@ from .atproto import ( resolve_pds_from_did, ) from .atproto.oauth import pds_authed_req -from .db import KV, close_db_connection, init_db +from .db import KV, close_db_connection, get_db, init_db from .oauth import get_auth_session, oauth, save_auth_session from .types import OAuthSession @@ -25,7 +26,7 @@ SCHEMA = "at.ligo" @app.before_request -def load_user_to_context(): +async def load_user_to_context(): g.user = get_auth_session(session) @@ -34,7 +35,7 @@ def get_user() -> OAuthSession | None: @app.teardown_appcontext -def app_teardown(exception: BaseException | None): +async def app_teardown(exception: BaseException | None): close_db_connection(exception) @@ -47,9 +48,13 @@ def page_home(): async def page_profile(atid: str): reload = request.args.get("reload") is not None + db = get_db(app) + didkv = KV(db, "did_from_handle") + pdskv = KV(db, "pds_from_did") + if atid.startswith("@"): handle = atid[1:].lower() - did = await resolve_did_from_handle(handle, reload=reload) + did = await resolve_did_from_handle(handle, kv=didkv, reload=reload) if did is None: return render_template("error.html", message="did not found"), 404 elif is_valid_did(atid): @@ -60,14 +65,14 @@ async def page_profile(atid: str): if _is_did_blocked(did): return render_template("error.html", message="profile not found"), 404 - kv = KV(app, "pds_from_did") - pds = await resolve_pds_from_did(did, kv, reload=reload) - if pds is None: - return render_template("error.html", message="pds not found"), 404 - (profile, _), links = await asyncio.gather( - load_profile(pds, did, reload=reload), - load_links(pds, did, reload=reload), - ) + async with ClientSession() as client: + pds = await resolve_pds_from_did(client, did=did, kv=pdskv, reload=reload) + if pds is None: + return render_template("error.html", message="pds not found"), 404 + (profile, _), links = await asyncio.gather( + load_profile(client, pds, did, reload=reload), + load_links(client, pds, did, reload=reload), + ) if links is None: return render_template("error.html", message="profile not found"), 404 @@ -112,10 +117,11 @@ async def page_editor(): pds: str = user.pds_url handle: str | None = user.handle - (profile, from_bluesky), links = await asyncio.gather( - load_profile(pds, did, reload=True), - load_links(pds, did, reload=True), - ) + async with ClientSession() as client: + (profile, from_bluesky), links = await asyncio.gather( + load_profile(client, pds, did, reload=True), + load_links(client, pds, did, reload=True), + ) return render_template( "editor.html", @@ -197,6 +203,7 @@ def page_terms(): async def load_links( + client: ClientSession, pds: str, did: str, reload: bool = False, @@ -208,7 +215,7 @@ async def load_links( app.logger.debug(f"returning cached links for {did}") return json.loads(recordstr)["links"] - record = await get_record(pds, did, f"{SCHEMA}.actor.links", "self") + record = await get_record(client, pds, did, f"{SCHEMA}.actor.links", "self") if record is None: return None @@ -218,6 +225,7 @@ async def load_links( async def load_profile( + client: ClientSession, pds: str, did: str, fallback_with_bluesky: bool = True, @@ -231,9 +239,9 @@ async def load_profile( return json.loads(recordstr), False from_bluesky = False - record = await get_record(pds, did, f"{SCHEMA}.actor.profile", "self") + record = await get_record(client, pds, did, f"{SCHEMA}.actor.profile", "self") if record is None and fallback_with_bluesky: - record = await get_record(pds, did, "app.bsky.actor.profile", "self") + record = await get_record(client, pds, did, "app.bsky.actor.profile", "self") from_bluesky = True if record is None: return None, False diff --git a/src/oauth.py b/src/oauth.py index a37abf4..d62e858 100644 --- a/src/oauth.py +++ b/src/oauth.py @@ -1,7 +1,8 @@ -from typing import NamedTuple +from aiohttp.client import ClientSession from authlib.jose import JsonWebKey, Key from flask import Blueprint, current_app, jsonify, redirect, request, session, url_for from flask.sessions import SessionMixin +from typing import NamedTuple from urllib.parse import urlencode import json @@ -33,10 +34,12 @@ async def oauth_start(): db = get_db(current_app) pdskv = KV(db, "authserver_from_pds") + client = ClientSession() + if is_valid_handle(username) or is_valid_did(username): login_hint = username kv = KV(db, "did_from_handle") - identity = await resolve_identity(username, didkv=kv) + identity = await resolve_identity(client, username, didkv=kv) if identity is None: return "couldnt resolve identity", 500 did, handle, doc = identity @@ -44,24 +47,28 @@ async def oauth_start(): if not pds_url: return "pds not found", 404 current_app.logger.debug(f"account PDS: {pds_url}") - authserver_url = await resolve_authserver_from_pds(pds_url, pdskv) + authserver_url = await resolve_authserver_from_pds(client, pds_url, pdskv) if not authserver_url: return "authserver not found", 404 elif username.startswith("https://") and is_safe_url(username): did, handle, pds_url = None, None, None login_hint = None - authserver_url = await resolve_authserver_from_pds(username, pdskv) or username + authserver_url = ( + await resolve_authserver_from_pds(client, username, pdskv) or username + ) else: return "not a valid handle, did or auth server", 400 current_app.logger.debug(f"Authserver: {authserver_url}") assert is_safe_url(authserver_url) - authserver_meta = await fetch_authserver_meta(authserver_url) + authserver_meta = await fetch_authserver_meta(client, authserver_url) if not authserver_meta: return "no authserver meta", 404 + await client.close() + # Auth dpop_private_jwk: Key = JsonWebKey.generate_key("EC", "P-256", is_private=True) scope = "atproto transition:generic" @@ -133,9 +140,12 @@ async def oauth_callback(): assert auth_request.authserver_iss == authserver_iss assert auth_request.state == state + client = ClientSession() + app_url = request.url_root.replace("http://", "https://") CLIENT_SECRET_JWK = JsonWebKey.import_key(current_app.config["CLIENT_SECRET_JWK"]) tokens, dpop_authserver_nonce = await initial_token_request( + client, auth_request, authorization_code, app_url, @@ -155,16 +165,22 @@ async def oauth_callback(): else: did = tokens.sub assert is_valid_did(did) - identity = await resolve_identity(did, didkv=didkv) + 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 - authserver_url = await resolve_authserver_from_pds(pds_url, authserverkv) + authserver_url = await resolve_authserver_from_pds( + client, + pds_url, + authserverkv, + ) assert authserver_url == authserver_iss + await client.close() + assert row.scope == tokens.scope assert pds_url is not None