diff --git a/.gitignore b/.gitignore index 9b92495..bc6b7f9 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,4 @@ .env .venv +*.db *.pyc diff --git a/src/atproto2/__init__.py b/src/atproto2/__init__.py index 07ba327..c978b53 100644 --- a/src/atproto2/__init__.py +++ b/src/atproto2/__init__.py @@ -149,6 +149,13 @@ def resolve_authserver_meta(authserver_url: str) -> dict[str, str] | None: return meta +def get_record(pds: str, repo: str, collection: str, record: str) -> str | None: + response = http_get( + f"{pds}/xrpc/com.atproto.repo.getRecord?repo={repo}&collection={collection}&rkey={record}" + ) + return response + + def http_get_json(url: str) -> Any | None: response = requests.get(url) if response.ok: diff --git a/src/db.py b/src/db.py new file mode 100644 index 0000000..5045ec6 --- /dev/null +++ b/src/db.py @@ -0,0 +1,26 @@ +import sqlite3 + +from flask import Flask, g + + +def get_db(app: Flask) -> sqlite3.Connection: + db: sqlite3.Connection | None = g.get("db", None) + if db is None: + db_path: str = app.config.get("DATABASE_URL", "ligoat.db") + db = g.db = sqlite3.connect(db_path) + 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() + + +def init_db(app: Flask): + with app.app_context(): + db = get_db(app) + with app.open_resource("schema.sql", mode="r") as schema: + _ = db.cursor().executescript(schema.read()) + db.commit() diff --git a/src/main.py b/src/main.py index 52f0d21..70dd2f0 100644 --- a/src/main.py +++ b/src/main.py @@ -1,26 +1,46 @@ from atproto import Client -from atproto.exceptions import AtProtocolError from atproto_client.models import ComAtprotoRepoCreateRecord -from atproto_client.models.app.bsky.actor.defs import ProfileViewDetailed -from flask import Flask, session, redirect, render_template, request +from flask import Flask, g, session, redirect, render_template, request, url_for from urllib import request as http_request import json -from .atproto2 import resolve_did_from_handle, resolve_pds_from_did +from .atproto2 import get_record, resolve_did_from_handle, resolve_pds_from_did +from .db import close_db_connection, get_db, init_db from .oauth import oauth app = Flask(__name__) _ = app.config.from_prefixed_env() app.register_blueprint(oauth) +init_db(app) -pdss: dict[str, str] = {} -dids: dict[str, str] = {} links: dict[str, list[dict[str, str]]] = {} profiles: dict[str, tuple[str, str]] = {} SCHEMA = "one.nauta" +@app.before_request +def load_user_to_context(): + did: str | None = session.get("user_did") + if did is None: + g.user = None + else: + db = get_db(app) + g.user = db.execute( + "select * from oauth_session where did = ?", + (did,), + ).fetchone() + + +def get_user() -> dict[str, str] | None: + return g.user + + +@app.teardown_appcontext +def app_teardown(exception: BaseException | None): + close_db_connection(exception) + + @app.get("/") def page_home(): return render_template("index.html") @@ -50,34 +70,36 @@ def page_profile(handle: str): @app.get("/login") def page_login(): - if "session" in session: + if get_user() is not None: return redirect("/editor") return render_template("login.html") +@app.post("/login") +def auth_login(): + username = request.form.get("username") + if not username: + return redirect(url_for("page_login"), 303) + return redirect(url_for("oauth.oauth_start", username=username), 303) + + @app.get("/editor") def page_editor(): - sess: str | None = session.get("session") - if sess is None or not sess: + user = get_user() + if user is None: return redirect("/login") - client = Client() - profile: ProfileViewDetailed | None - try: - profile = client.login(session_string=sess) - except AtProtocolError: - session.clear() - return redirect("/login", 303) - pds = resolve_pds_from_did(profile.did) - if not pds: - return "did not found", 404 - pro, from_bluesky = load_profile(pds, profile.did, reload=True) - links = load_links(pds, profile.did, reload=True) or [{"background": "#fa0"}] + did: str = user["did"] + pds: str = user["pds_url"] + handle: str | None = user["handle"] + + profile, from_bluesky = load_profile(pds, did, reload=True) + links = load_links(pds, did, reload=True) or [{"background": "#fa0"}] return render_template( "editor.html", - handle=profile.handle, - profile=pro, + handle=handle, + profile=profile, profile_from_bluesky=from_bluesky, links=json.dumps(links), ) @@ -85,11 +107,12 @@ def page_editor(): @app.post("/editor/profile") def post_editor_profile(): - sess: str | None = session.get("session") - if sess is None or not sess: + user = get_user() + if user is None: return redirect("/login", 303) + client = Client() - profile = client.login(session_string=sess) + profile = client.login(session_string=user["did"]) display_name = request.form.get("displayName") description = request.form.get("description") or "" @@ -190,13 +213,6 @@ def load_profile( return profile, from_bluesky -def get_record(pds: str, repo: str, collection: str, record: str) -> str | None: - response = http_get( - f"{pds}/xrpc/com.atproto.repo.getRecord?repo={repo}&collection={collection}&rkey={record}" - ) - return response - - def put_record(client: Client, repo: str, collection: str, rkey: str, record): data_model = ComAtprotoRepoCreateRecord.Data( collection=collection, @@ -223,15 +239,12 @@ def http_get(url: str) -> str | None: @app.route("/auth/logout") def auth_logout(): + user = get_user() + if user is not None: + db = get_db(app) + cursor = db.cursor() + _ = cursor.execute("delete from oauth_session where did = ?", (user["did"],)) + db.commit() + cursor.close() session.clear() - return redirect("/") - - -@app.post("/auth/login") -def auth_login(): - handle = request.form.get("handle") - if not handle: - return redirect("/login", 303) - if handle.startswith("@"): - handle = handle[1:] - return redirect(app.url_for("oauth.oauth_start", username=handle)) + return redirect("/", 303) diff --git a/src/oauth.py b/src/oauth.py index d945535..366db2d 100644 --- a/src/oauth.py +++ b/src/oauth.py @@ -5,28 +5,19 @@ from urllib.parse import urlencode import json from .atproto2.atproto_oauth import initial_token_request, send_par_auth_request - from .atproto2.atproto_security import is_safe_url - from .atproto2 import ( pds_endpoint_from_doc, resolve_authserver_from_pds, resolve_authserver_meta, resolve_identity, ) +from .db import get_db oauth = Blueprint("oauth", __name__, url_prefix="/oauth") oauth_auth_requests: dict[str, dict[str, str]] = {} -oauth_session: dict[str, dict[str, str]] = {} - - -@oauth.get("/home") -def oauth_home(): - user_did = session["user_did"] - user_handle = session["user_handle"] - return f"{user_did} {user_handle}" @oauth.get("/start") @@ -75,7 +66,7 @@ def oauth_start(): dpop_private_jwk, ) if resp.status_code == 400: - print(f"PAR HTTP 400: {resp.json()}") + current_app.logger.info(f"PAR HTTP 400: {resp.json()}") resp.raise_for_status() par_request_uri = resp.json()["request_uri"] @@ -105,7 +96,7 @@ def oauth_callback(): auth_request = oauth_auth_requests.get(state) if auth_request is None: - return redirect(url_for("oauth.oauth_home"), 303) + return redirect(url_for("page_login"), 303) current_app.logger.debug(f"Deleting auth request for state={state}") _ = oauth_auth_requests.pop(state) @@ -135,23 +126,30 @@ def oauth_callback(): assert row["scope"] == tokens["scope"] - oauth_session[did] = { - "did": did, - "handle": handle, - "pds_url": pds_url, - "authserver_iss": authserver_iss, - "access_token": tokens["access_token"], - "refresh_token": tokens["refresh_token"], - "dpop_authserver_nonce": dpop_authserver_nonce, - "dpop_private_jwk": auth_request["dpop_private_jwk"], - } - current_app.logger.debug("storing user did and handle") + db = get_db(current_app) + cursor = db.cursor() + _ = cursor.execute( + "insert or replace into oauth_session values (?, ?, ?, ?, ?, ?, ?, ?, ?)", + ( + did, + handle, + pds_url, + authserver_iss, + tokens["access_token"], + tokens["refresh_token"], + dpop_authserver_nonce, + None, + auth_request["dpop_private_jwk"], + ), + ) + db.commit() + cursor.close() session["user_did"] = did session["user_handle"] = auth_request["handle"] - return redirect(url_for("oauth.oauth_home")) + return redirect(url_for("page_login")) @oauth.get("/metadata") diff --git a/src/schema.sql b/src/schema.sql new file mode 100644 index 0000000..ca44f9d --- /dev/null +++ b/src/schema.sql @@ -0,0 +1,11 @@ +create table if not exists oauth_session ( + did text not null primary key, + handle text, + pds_url text not null, + authserver_iss text not null, + access_token text, + refresh_token text, + dpop_authserver_nonce text not null, + dpop_pds_nonce text, + dpop_private_jwk text not null +) strict, without rowid; diff --git a/src/templates/editor.html b/src/templates/editor.html index a882d19..f1174c9 100644 --- a/src/templates/editor.html +++ b/src/templates/editor.html @@ -19,7 +19,7 @@

see profile ยท - logout + logout

profile

diff --git a/src/templates/login.html b/src/templates/login.html index 2d4875c..7adefc8 100644 --- a/src/templates/login.html +++ b/src/templates/login.html @@ -13,10 +13,10 @@

atlinks

log in to your account -
+