diff --git a/add_links/api.py b/add_links/api.py index 4b9abc5..6bd7577 100644 --- a/add_links/api.py +++ b/add_links/api.py @@ -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" diff --git a/test_api.py b/test_api.py index abfd292..df1ccf5 100644 --- a/test_api.py +++ b/test_api.py @@ -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() diff --git a/test_web_view.py b/test_web_view.py index ba2392b..da460d2 100644 --- a/test_web_view.py +++ b/test_web_view.py @@ -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")] diff --git a/web_view.py b/web_view.py index 85a9ce4..00866cd 100755 --- a/web_view.py +++ b/web_view.py @@ -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