diff --git a/src/api/auth.py b/src/api/auth.py new file mode 100644 index 0000000..1985684 --- /dev/null +++ b/src/api/auth.py @@ -0,0 +1,32 @@ +import os +from functools import wraps + +from flask import abort, g, jsonify, make_response, request + +from src.api import tokens_db + + +def _abort_401(message: str): + resp = make_response(jsonify({"message": message}), 401) + abort(resp) + + +def require_token(fn): + @wraps(fn) + def wrapper(*args, **kwargs): + header = request.headers.get("Authorization", "") + if not header.startswith("Bearer "): + _abort_401("missing_token") + token = header[len("Bearer ") :].strip() + if not token: + _abort_401("missing_token") + db_path = os.environ["USERS_DB_PATH"] + row = tokens_db.get_token_by_plaintext(db_path, token) + if row is None: + _abort_401("invalid_token") + if row["revoked_at"] is not None: + _abort_401("revoked_token") + g.token_id = row["id"] + return fn(*args, **kwargs) + + return wrapper diff --git a/tests/api/test_auth.py b/tests/api/test_auth.py new file mode 100644 index 0000000..1584bed --- /dev/null +++ b/tests/api/test_auth.py @@ -0,0 +1,59 @@ +from flask import Flask, g, jsonify + +from src.api import tokens_db +from src.api.auth import require_token + + +def _make_app(): + app = Flask(__name__) + + @app.route("/protected") + @require_token + def protected(): + return jsonify({"token_id": g.token_id}) + + return app + + +def test_missing_header_returns_401(temp_db): + app = _make_app() + resp = app.test_client().get("/protected") + assert resp.status_code == 401 + assert resp.get_json()["message"] == "missing_token" + + +def test_bearer_without_value_returns_401(temp_db): + app = _make_app() + resp = app.test_client().get("/protected", headers={"Authorization": "Bearer "}) + assert resp.status_code == 401 + assert resp.get_json()["message"] == "missing_token" + + +def test_invalid_token_returns_401(temp_db): + app = _make_app() + resp = app.test_client().get( + "/protected", headers={"Authorization": "Bearer decpinfo_unknown"} + ) + assert resp.status_code == 401 + assert resp.get_json()["message"] == "invalid_token" + + +def test_revoked_token_returns_401(temp_db): + token, token_id = tokens_db.create_token(temp_db, "x") + tokens_db.revoke_token(temp_db, token_id) + app = _make_app() + resp = app.test_client().get( + "/protected", headers={"Authorization": f"Bearer {token}"} + ) + assert resp.status_code == 401 + assert resp.get_json()["message"] == "revoked_token" + + +def test_valid_token_sets_g_and_calls_view(temp_db): + token, token_id = tokens_db.create_token(temp_db, "x") + app = _make_app() + resp = app.test_client().get( + "/protected", headers={"Authorization": f"Bearer {token}"} + ) + assert resp.status_code == 200 + assert resp.get_json()["token_id"] == token_id