Use OAuth sessions for web API requests
This commit is contained in:
parent
e022872767
commit
118ec9fc9b
4 changed files with 125 additions and 7 deletions
|
|
@ -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]:
|
def _get_active_session() -> tuple[requests.sessions.Session, str]:
|
||||||
"""Return active session and auth mode."""
|
"""Return active session and auth mode."""
|
||||||
try:
|
if oauth_session := _get_request_oauth_session():
|
||||||
from flask import g
|
return oauth_session, "oauth"
|
||||||
if hasattr(g, "oauth_session") and g.oauth_session is not None:
|
|
||||||
return g.oauth_session, "oauth" # type: ignore[return-value]
|
|
||||||
except RuntimeError:
|
|
||||||
pass
|
|
||||||
print("WARNING: using unauthenticated session", file=sys.stderr)
|
print("WARNING: using unauthenticated session", file=sys.stderr)
|
||||||
return get_session(), "anon"
|
return get_session(), "anon"
|
||||||
|
|
||||||
|
|
|
||||||
14
test_api.py
14
test_api.py
|
|
@ -3,6 +3,7 @@ from unittest.mock import Mock, patch
|
||||||
|
|
||||||
from simplejson.scanner import JSONDecodeError
|
from simplejson.scanner import JSONDecodeError
|
||||||
|
|
||||||
|
import web_view
|
||||||
from add_links import api
|
from add_links import api
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -25,6 +26,19 @@ class ApiGetTests(unittest.TestCase):
|
||||||
|
|
||||||
self.assertEqual(str(ctx.exception), response.text)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|
|
||||||
|
|
@ -37,6 +37,8 @@ class ArticlePageSkipTests(unittest.TestCase):
|
||||||
}
|
}
|
||||||
|
|
||||||
with (
|
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.api.get_wiki_info", return_value=None),
|
||||||
patch("web_view.search_count", return_value=3),
|
patch("web_view.search_count", return_value=3),
|
||||||
patch("web_view.search_count_with_link", return_value=0),
|
patch("web_view.search_count_with_link", return_value=0),
|
||||||
|
|
@ -62,6 +64,8 @@ class ArticlePageSkipTests(unittest.TestCase):
|
||||||
]
|
]
|
||||||
|
|
||||||
with (
|
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.api.get_wiki_info", return_value="Documentary film"),
|
||||||
patch("web_view.search_count", return_value=2),
|
patch("web_view.search_count", return_value=2),
|
||||||
patch("web_view.search_count_with_link", return_value=0),
|
patch("web_view.search_count_with_link", return_value=0),
|
||||||
|
|
@ -78,6 +82,8 @@ class ArticlePageSkipTests(unittest.TestCase):
|
||||||
hits = [hit("documentary film")]
|
hits = [hit("documentary film")]
|
||||||
|
|
||||||
with (
|
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.api.get_wiki_info", return_value="Documentary film"),
|
||||||
patch("web_view.search_count", return_value=1),
|
patch("web_view.search_count", return_value=1),
|
||||||
patch("web_view.search_count_with_link", return_value=0),
|
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()
|
self.client = web_view.app.test_client()
|
||||||
|
|
||||||
def test_redirect_target_is_not_valid_hit(self) -> None:
|
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(
|
response = self.client.get(
|
||||||
"/api/1/valid_hit",
|
"/api/1/valid_hit",
|
||||||
query_string={
|
query_string={
|
||||||
|
|
@ -111,6 +121,53 @@ class ValidHitTests(unittest.TestCase):
|
||||||
get_diff.assert_not_called()
|
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):
|
class ProtectedCandidateTests(unittest.TestCase):
|
||||||
def test_edit_protected_hits_are_filtered_out(self) -> None:
|
def test_edit_protected_hits_are_filtered_out(self) -> None:
|
||||||
hits = [hit("Editable"), hit("Protected")]
|
hits = [hit("Editable"), hit("Protected")]
|
||||||
|
|
|
||||||
29
web_view.py
29
web_view.py
|
|
@ -207,6 +207,22 @@ def global_user() -> None:
|
||||||
flask.g.oauth_session = mediawiki_oauth.get_oauth_session()
|
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("/")
|
@app.route("/")
|
||||||
def index() -> str | Response:
|
def index() -> str | Response:
|
||||||
"""Index page."""
|
"""Index page."""
|
||||||
|
|
@ -406,6 +422,10 @@ def _record_skip(from_title: str, hit_title: str) -> None:
|
||||||
|
|
||||||
def handle_post(url_title: str) -> Response:
|
def handle_post(url_title: str) -> Response:
|
||||||
"""Handle POST request."""
|
"""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()
|
from_title = url_title.replace("_", " ").strip()
|
||||||
hit_title = flask.request.form["hit"]
|
hit_title = flask.request.form["hit"]
|
||||||
try:
|
try:
|
||||||
|
|
@ -436,6 +456,9 @@ def article_page(url_title: str) -> str | Response:
|
||||||
if flask.request.method == "POST":
|
if flask.request.method == "POST":
|
||||||
return handle_post(url_title)
|
return handle_post(url_title)
|
||||||
|
|
||||||
|
if not is_logged_in():
|
||||||
|
return login_redirect()
|
||||||
|
|
||||||
from_title = url_title.replace("_", " ").strip()
|
from_title = url_title.replace("_", " ").strip()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
@ -525,6 +548,9 @@ def save_done() -> str:
|
||||||
@app.route("/api/1/hits")
|
@app.route("/api/1/hits")
|
||||||
def api_hits() -> werkzeug.wrappers.response.Response:
|
def api_hits() -> werkzeug.wrappers.response.Response:
|
||||||
"""Return candidates for the given article title."""
|
"""Return candidates for the given article title."""
|
||||||
|
if not is_logged_in():
|
||||||
|
return login_required_json()
|
||||||
|
|
||||||
title = flask.request.args.get("title")
|
title = flask.request.args.get("title")
|
||||||
assert title
|
assert title
|
||||||
ret = core.do_search(title)
|
ret = core.do_search(title)
|
||||||
|
|
@ -537,6 +563,9 @@ def api_hits() -> werkzeug.wrappers.response.Response:
|
||||||
@app.route("/api/1/valid_hit")
|
@app.route("/api/1/valid_hit")
|
||||||
def api_valid_hit() -> werkzeug.wrappers.response.Response:
|
def api_valid_hit() -> werkzeug.wrappers.response.Response:
|
||||||
"""Check if a candidate article has a valid unlinked mention."""
|
"""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_to = flask.request.args["link_to"]
|
||||||
link_from = flask.request.args["link_from"]
|
link_from = flask.request.args["link_from"]
|
||||||
redirect_to = flask.request.args.get("redirect_to") or None
|
redirect_to = flask.request.args.get("redirect_to") or None
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue