Use OAuth sessions for web API requests

This commit is contained in:
Edward Betts 2026-07-22 10:58:10 +01:00
parent e022872767
commit 118ec9fc9b
4 changed files with 125 additions and 7 deletions

View file

@ -85,14 +85,32 @@ webpage_error = (
)
def _get_request_oauth_session() -> requests.sessions.Session | None:
"""Return the logged-in user's OAuth session in a Flask request, if present."""
try:
from flask import g, has_request_context
if not has_request_context():
return None
oauth_session = getattr(g, "oauth_session", None)
if oauth_session is not None:
return typing.cast(requests.sessions.Session, oauth_session)
from . import mediawiki_oauth
return typing.cast(
requests.sessions.Session | None, mediawiki_oauth.get_oauth_session()
)
except RuntimeError:
return None
def _get_active_session() -> tuple[requests.sessions.Session, str]:
"""Return active session and auth mode."""
try:
from flask import g
if hasattr(g, "oauth_session") and g.oauth_session is not None:
return g.oauth_session, "oauth" # type: ignore[return-value]
except RuntimeError:
pass
if oauth_session := _get_request_oauth_session():
return oauth_session, "oauth"
print("WARNING: using unauthenticated session", file=sys.stderr)
return get_session(), "anon"

View file

@ -3,6 +3,7 @@ from unittest.mock import Mock, patch
from simplejson.scanner import JSONDecodeError
import web_view
from add_links import api
@ -25,6 +26,19 @@ class ApiGetTests(unittest.TestCase):
self.assertEqual(str(ctx.exception), response.text)
def test_active_session_uses_logged_in_request_session(self) -> None:
oauth_session = Mock()
with web_view.app.test_request_context("/"):
with patch(
"add_links.mediawiki_oauth.get_oauth_session",
return_value=oauth_session,
):
session, auth_mode = api._get_active_session()
self.assertIs(session, oauth_session)
self.assertEqual(auth_mode, "oauth")
if __name__ == "__main__":
unittest.main()

View file

@ -37,6 +37,8 @@ class ArticlePageSkipTests(unittest.TestCase):
}
with (
patch("web_view.mediawiki_oauth.get_username", return_value="Test user"),
patch("web_view.mediawiki_oauth.get_oauth_session", return_value=object()),
patch("web_view.api.get_wiki_info", return_value=None),
patch("web_view.search_count", return_value=3),
patch("web_view.search_count_with_link", return_value=0),
@ -62,6 +64,8 @@ class ArticlePageSkipTests(unittest.TestCase):
]
with (
patch("web_view.mediawiki_oauth.get_username", return_value="Test user"),
patch("web_view.mediawiki_oauth.get_oauth_session", return_value=object()),
patch("web_view.api.get_wiki_info", return_value="Documentary film"),
patch("web_view.search_count", return_value=2),
patch("web_view.search_count_with_link", return_value=0),
@ -78,6 +82,8 @@ class ArticlePageSkipTests(unittest.TestCase):
hits = [hit("documentary film")]
with (
patch("web_view.mediawiki_oauth.get_username", return_value="Test user"),
patch("web_view.mediawiki_oauth.get_oauth_session", return_value=object()),
patch("web_view.api.get_wiki_info", return_value="Documentary film"),
patch("web_view.search_count", return_value=1),
patch("web_view.search_count_with_link", return_value=0),
@ -96,7 +102,11 @@ class ValidHitTests(unittest.TestCase):
self.client = web_view.app.test_client()
def test_redirect_target_is_not_valid_hit(self) -> None:
with patch("web_view.get_diff") as get_diff:
with (
patch("web_view.mediawiki_oauth.get_username", return_value="Test user"),
patch("web_view.mediawiki_oauth.get_oauth_session", return_value=object()),
patch("web_view.get_diff") as get_diff,
):
response = self.client.get(
"/api/1/valid_hit",
query_string={
@ -111,6 +121,53 @@ class ValidHitTests(unittest.TestCase):
get_diff.assert_not_called()
class LoginRequiredTests(unittest.TestCase):
def setUp(self) -> None:
web_view.app.config["TESTING"] = True
self.client = web_view.app.test_client()
def test_article_page_redirects_to_login_before_api_calls(self) -> None:
with (
patch("web_view.mediawiki_oauth.get_username", return_value=None),
patch("web_view.mediawiki_oauth.get_oauth_session", return_value=None),
patch("web_view.api.get_wiki_info") as get_wiki_info,
):
response = self.client.get("/link/DVD_documentary")
self.assertEqual(response.status_code, 302)
self.assertIn("/oauth/start", response.location)
get_wiki_info.assert_not_called()
def test_valid_hit_requires_login_before_diff_call(self) -> None:
with (
patch("web_view.mediawiki_oauth.get_username", return_value=None),
patch("web_view.mediawiki_oauth.get_oauth_session", return_value=None),
patch("web_view.get_diff") as get_diff,
):
response = self.client.get(
"/api/1/valid_hit",
query_string={"link_to": "Target", "link_from": "Candidate"},
)
self.assertEqual(response.status_code, 401)
self.assertEqual(response.get_json(), {"error": "login required"})
get_diff.assert_not_called()
def test_hits_api_requires_login_before_search(self) -> None:
with (
patch("web_view.mediawiki_oauth.get_username", return_value=None),
patch("web_view.mediawiki_oauth.get_oauth_session", return_value=None),
patch("web_view.core.do_search") as do_search,
):
response = self.client.get(
"/api/1/hits", query_string={"title": "Target"}
)
self.assertEqual(response.status_code, 401)
self.assertEqual(response.get_json(), {"error": "login required"})
do_search.assert_not_called()
class ProtectedCandidateTests(unittest.TestCase):
def test_edit_protected_hits_are_filtered_out(self) -> None:
hits = [hit("Editable"), hit("Protected")]

View file

@ -207,6 +207,22 @@ def global_user() -> None:
flask.g.oauth_session = mediawiki_oauth.get_oauth_session()
def is_logged_in() -> bool:
"""Return true when the current request has a Wikipedia OAuth session."""
return flask.g.oauth_session is not None
def login_redirect() -> Response:
"""Redirect the current browser request through Wikipedia OAuth."""
next_url = flask.request.full_path.rstrip("?")
return flask.redirect(flask.url_for("start_oauth", next=next_url))
def login_required_json() -> tuple[werkzeug.wrappers.response.Response, int]:
"""Return a JSON login-required response without making MediaWiki API calls."""
return flask.jsonify(error="login required"), 401
@app.route("/")
def index() -> str | Response:
"""Index page."""
@ -406,6 +422,10 @@ def _record_skip(from_title: str, hit_title: str) -> None:
def handle_post(url_title: str) -> Response:
"""Handle POST request."""
if not is_logged_in():
next_url = flask.url_for("article_page", url_title=url_title)
return flask.redirect(flask.url_for("start_oauth", next=next_url))
from_title = url_title.replace("_", " ").strip()
hit_title = flask.request.form["hit"]
try:
@ -436,6 +456,9 @@ def article_page(url_title: str) -> str | Response:
if flask.request.method == "POST":
return handle_post(url_title)
if not is_logged_in():
return login_redirect()
from_title = url_title.replace("_", " ").strip()
try:
@ -525,6 +548,9 @@ def save_done() -> str:
@app.route("/api/1/hits")
def api_hits() -> werkzeug.wrappers.response.Response:
"""Return candidates for the given article title."""
if not is_logged_in():
return login_required_json()
title = flask.request.args.get("title")
assert title
ret = core.do_search(title)
@ -537,6 +563,9 @@ def api_hits() -> werkzeug.wrappers.response.Response:
@app.route("/api/1/valid_hit")
def api_valid_hit() -> werkzeug.wrappers.response.Response:
"""Check if a candidate article has a valid unlinked mention."""
if not is_logged_in():
return login_required_json()
link_to = flask.request.args["link_to"]
link_from = flask.request.args["link_from"]
redirect_to = flask.request.args.get("redirect_to") or None