diff --git a/src/atproto/__init__.py b/src/atproto/__init__.py index 3b4e6c9..e676c66 100644 --- a/src/atproto/__init__.py +++ b/src/atproto/__init__.py @@ -2,6 +2,7 @@ import asyncio from os import getenv from re import match as regex_match from typing import Any, TypeGuard +from urllib.parse import urljoin from aiodns import DNSResolver from aiodns import error as dns_error @@ -96,8 +97,10 @@ async def resolve_identity_microcosm( 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}" + url = urljoin( + "https://slingshot.microcosm.blue", + f"/xrpc/com.bad-example.identity.resolveMiniDoc?identifier={query}", + ) response = await client.get(url) if not response.ok: return None @@ -241,7 +244,7 @@ async def resolve_authserver_from_pds( return authserver_url assert is_safe_url(pds_url) - endpoint = f"{pds_url}/.well-known/oauth-protected-resource" + endpoint = urljoin(pds_url, "/.well-known/oauth-protected-resource") response = await client.get(endpoint) if response.status != 200: return None @@ -258,7 +261,7 @@ async def fetch_authserver_meta( """Returns metadata from the authserver""" assert is_safe_url(authserver_url) - endpoint = f"{authserver_url}/.well-known/oauth-authorization-server" + endpoint = urljoin(authserver_url, "/.well-known/oauth-authorization-server") response = await client.get(endpoint) if not response.ok: return None @@ -278,7 +281,8 @@ async def get_record( """Retrieve record from PDS. Verifies type is the same as collection name.""" params = {"repo": repo, "collection": collection, "rkey": record} - response = await client.get(f"{pds}/xrpc/com.atproto.repo.getRecord", params=params) + url = urljoin(pds, "/xrpc/com.atproto.repo.getRecord") + response = await client.get(url, params=params) if not response.ok: return None parsed = await response.json() diff --git a/src/config.py b/src/config.py new file mode 100644 index 0000000..380240a --- /dev/null +++ b/src/config.py @@ -0,0 +1,26 @@ +from sqlite3 import Connection +from typing import NamedTuple + +from flask import Flask + +from src.db import get_db + + +class AuthServer(NamedTuple): + name: str + url: str + + +class Config: + db: Connection + + def __init__(self, app: Flask): + self.db = get_db(app, name="config") + + def auth_servers(self) -> list[AuthServer]: + raw = ( + self.db.cursor() + .execute("select name, url from pdss order by relevance desc") + .fetchall() + ) + return [AuthServer(*r) for r in raw] diff --git a/src/config.sql b/src/config.sql new file mode 100644 index 0000000..dd3840b --- /dev/null +++ b/src/config.sql @@ -0,0 +1,7 @@ +create table if not exists pdss ( + name text not null unique, + url text not null unique, + relevance integer not null +) strict; + +create index if not exists pdss_by_relevance on pdss(relevance desc); diff --git a/src/db.py b/src/db.py index 985939e..e877467 100644 --- a/src/db.py +++ b/src/db.py @@ -1,7 +1,7 @@ import sqlite3 from logging import Logger from sqlite3 import Connection -from typing import Generic, cast, override +from typing import Generic, Literal, cast, override from flask import Flask, g @@ -15,7 +15,7 @@ class KV(BaseKV, Generic[K, V]): prefix: str def __init__(self, app: Connection | Flask, logger: Logger, prefix: str): - self.db = app if isinstance(app, Connection) else get_db(app) + self.db = app if isinstance(app, Connection) else get_db(app, name="keyval") self.logger = logger self.prefix = prefix @@ -42,25 +42,31 @@ class KV(BaseKV, Generic[K, V]): self.db.commit() -def get_db(app: Flask) -> sqlite3.Connection: - db: sqlite3.Connection | None = g.get("db", None) +type DatabaseName = Literal["config"] | Literal["keyval"] + + +def get_db(app: Flask, name: DatabaseName) -> sqlite3.Connection: + global_key = f"{name}_db" + db: sqlite3.Connection | None = g.get(global_key, None) if db is None: - db_path: str = app.config.get("DATABASE_URL", "ligoat.db") - db = g.db = sqlite3.connect(db_path, check_same_thread=False) + db_path: str = app.config[f"{name.upper()}_DB_URL"] + db = sqlite3.connect(db_path, check_same_thread=False) + setattr(g, global_key, db) # return rows as dict-like objects db.row_factory = sqlite3.Row return db def close_db_connection(_exception: BaseException | None): - db: sqlite3.Connection | None = g.get("db", None) - if db is not None: - db.close() + for name in ["keyval", "config"]: + db: sqlite3.Connection | None = g.pop(f"{name}_db", None) + if db is not None: + db.close() -def init_db(app: Flask): +def init_db(app: Flask, name: DatabaseName) -> None: with app.app_context(): - db = get_db(app) - with app.open_resource("schema.sql", mode="r") as schema: + db = get_db(app, name) + with app.open_resource(f"{name}.sql", mode="r") as schema: _ = db.cursor().executescript(schema.read()) db.commit() diff --git a/src/schema.sql b/src/keyval.sql similarity index 100% rename from src/schema.sql rename to src/keyval.sql diff --git a/src/main.py b/src/main.py index 56d9017..4893f55 100644 --- a/src/main.py +++ b/src/main.py @@ -20,6 +20,7 @@ from src.auth import ( refresh_auth_session, save_auth_session, ) +from src.config import AuthServer, Config from src.db import KV, close_db_connection, get_db, init_db from src.oauth import oauth @@ -28,7 +29,8 @@ _ = app.config.from_prefixed_env() app.register_blueprint(oauth) htmx = HTMX() htmx.init_app(app) -init_db(app) +init_db(app, name="config") +init_db(app, name="keyval") @app.before_request @@ -58,7 +60,7 @@ def page_home(): async def page_profile(atid: str): reload = request.args.get("reload") is not None - db = get_db(app) + db = get_db(app, name="keyval") didkv = KV[Handle, DID](db, app.logger, "did_from_handle") pdskv = KV[DID, PdsUrl](db, app.logger, "pds_from_did") @@ -103,28 +105,14 @@ async def page_profile(atid: str): ) -class AuthServer(NamedTuple): - name: str - url: str - - -auth_servers: list[AuthServer] = [ - AuthServer("Bluesky", "https://bsky.social"), - AuthServer("Blacksky", "https://blacksky.app"), - AuthServer("Northsky", "https://northsky.social"), - AuthServer("tangled.org", "https://tngl.sh"), - AuthServer("Witchraft Systems", "https://pds.witchcraft.systems"), - AuthServer("selfhosted.social", "https://selfhosted.social"), -] - -if app.debug: - auth_servers.append(AuthServer("pds.rip", "https://pds.rip")) - - @app.get("/login") async def page_login(): if await get_user() is not None: return redirect("/editor") + config = Config(app) + auth_servers = config.auth_servers() + if app.debug: + auth_servers.append(AuthServer("pds.rip", "https://pds.rip")) return render_template("login.html", auth_servers=auth_servers) diff --git a/src/oauth.py b/src/oauth.py index cc81d06..b529ad1 100644 --- a/src/oauth.py +++ b/src/oauth.py @@ -43,7 +43,7 @@ async def oauth_start(): if not username: return redirect(url_for("page_login"), 303) - db = get_db(current_app) + db = get_db(current_app, name="keyval") 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]( @@ -177,7 +177,7 @@ async def oauth_callback(): row = auth_request - db = get_db(current_app) + db = get_db(current_app, name="keyval") didkv = KV(db, current_app.logger, "did_from_handle") authserverkv = KV(db, current_app.logger, "authserver_from_pds") diff --git a/src/templates/login.html b/src/templates/login.html index d0badcd..f8bc40a 100644 --- a/src/templates/login.html +++ b/src/templates/login.html @@ -28,6 +28,7 @@ + {% if auth_servers %}
If you're unsure you can log in with...
@@ -36,6 +37,7 @@ {% endfor %}
+ {% endif %}