diff options
| -rw-r--r-- | tests/forms/test_forms.py (renamed from tests/test_forms.py) | 69 | ||||
| -rw-r--r-- | tests/forms/test_validators.py | 114 | ||||
| -rw-r--r-- | tests/tasks/test_extension.py | 175 | ||||
| -rw-r--r-- | tests/tasks/test_receiver.py | 5 | ||||
| -rw-r--r-- | tests/tasks/test_scanner.py | 97 | ||||
| -rw-r--r-- | tests/tasks/test_sender.py | 261 | ||||
| -rw-r--r-- | tests/test_models.py | 148 | ||||
| -rw-r--r-- | tests/test_url_security.py | 29 | ||||
| -rw-r--r-- | tests/test_views.py | 117 | ||||
| -rw-r--r-- | webmentions_ssg/tasks/extension.py | 102 | ||||
| -rw-r--r-- | webmentions_ssg/tasks/scanner.py | 1 |
11 files changed, 804 insertions, 314 deletions
diff --git a/tests/test_forms.py b/tests/forms/test_forms.py index fc16858..6fe1849 100644 --- a/tests/test_forms.py +++ b/tests/forms/test_forms.py @@ -2,19 +2,25 @@ # # SPDX-License-Identifier: BSD-3-Clause -from unittest.mock import patch +from unittest.mock import Mock import pytest from flask import Flask from werkzeug.datastructures import MultiDict -from webmentions_ssg.forms import EndpointForm +from webmentions_ssg.forms import AdminActionForm, EndpointForm, LoginForm VALID_SOURCE = "https://source.example/post" VALID_TARGET = "https://dennisfink.me/blog/example/" -@patch("webmentions_ssg.forms.validators.is_public_url", new=lambda url: True) +@pytest.fixture(autouse=True) +def public_urls(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + "webmentions_ssg.forms.validators.is_public_url", Mock(return_value=True) + ) + + @pytest.mark.parametrize( ("form_data", "invalid_field", "expected_error"), [ @@ -37,12 +43,6 @@ VALID_TARGET = "https://dennisfink.me/blog/example/" id="source-scheme", ), pytest.param( - {"source": VALID_TARGET, "target": VALID_TARGET}, - "source", - None, - id="source-not-equal-to-target", - ), - pytest.param( {"source": VALID_SOURCE}, "target", "This field is required.", @@ -60,31 +60,18 @@ VALID_TARGET = "https://dennisfink.me/blog/example/" "target must begin with http or https", id="target-scheme", ), - pytest.param( - {"source": VALID_SOURCE, "target": "https://example.com/post"}, - "target", - None, - id="target-allowed-hostname", - ), ], ) -def test_endpoint_form_rejects_invalid_data( - app: Flask, - form_data: dict[str, str], - invalid_field: str, - expected_error: str | None, +def test_endpoint_form_rejects_invalid_field_syntax( + app: Flask, form_data: dict[str, str], invalid_field: str, expected_error: str ) -> None: with app.test_request_context("/endpoint", method="POST"): form = EndpointForm(formdata=MultiDict(form_data), meta={"csrf": False}) assert not form.validate() - assert invalid_field in form.errors - - if expected_error is not None: - assert expected_error in form.errors[invalid_field] + assert expected_error in form.errors[invalid_field] -@patch("webmentions_ssg.forms.validators.is_public_url", new=lambda url: True) def test_endpoint_form_accepts_valid_data(app: Flask) -> None: with app.test_request_context("/endpoint", method="POST"): form = EndpointForm( @@ -94,3 +81,35 @@ def test_endpoint_form_accepts_valid_data(app: Flask) -> None: assert form.validate() assert form.errors == {} + + +@pytest.mark.parametrize( + ("form_data", "invalid_field"), + [ + pytest.param({"password": "secret"}, "username", id="username-required"), + pytest.param({"username": "admin"}, "password", id="password-required"), + ], +) +def test_login_form_requires_credentials( + app: Flask, form_data: dict[str, str], invalid_field: str +) -> None: + with app.test_request_context("/login", method="POST"): + form = LoginForm(formdata=MultiDict(form_data), meta={"csrf": False}) + + assert not form.validate() + assert "This field is required." in form.errors[invalid_field] + + +def test_login_form_accepts_credentials(app: Flask) -> None: + with app.test_request_context("/login", method="POST"): + form = LoginForm( + formdata=MultiDict({"username": "admin", "password": "secret"}), + meta={"csrf": False}, + ) + assert form.validate() + + +def test_admin_action_form_accepts_submission(app: Flask) -> None: + with app.test_request_context(method="POST"): + form = AdminActionForm(meta={"csrf": False}) + assert form.validate() diff --git a/tests/forms/test_validators.py b/tests/forms/test_validators.py new file mode 100644 index 0000000..c63ee54 --- /dev/null +++ b/tests/forms/test_validators.py @@ -0,0 +1,114 @@ +# SPDX-FileCopyrightText: 2026 Dennis Fink <me+coding@dennisfink.me> +# +# SPDX-License-Identifier: BSD-3-Clause + +from unittest.mock import Mock + +import pytest +from flask import Flask +from werkzeug.datastructures import MultiDict +from wtforms import Form, StringField, ValidationError + +from webmentions_ssg.forms.validators import AllowedHostname, NotEqualTo, PublicURL +from webmentions_ssg.url_security import AddressResolutionError + + +class ComparisonForm(Form): + source = StringField("Source") + target = StringField("Target") + + +class URLForm(Form): + url = StringField("URL") + + +def test_not_equal_to_accepts_different_values() -> None: + form = ComparisonForm(MultiDict({"source": "source", "target": "target"})) + NotEqualTo("target")(form, form.source) + + +def test_not_equal_to_rejects_equal_values() -> None: + form = ComparisonForm(MultiDict({"source": "same", "target": "same"})) + with pytest.raises(ValidationError, match="Field must not be equal to target"): + NotEqualTo("target")(form, form.source) + + +def test_not_equal_to_uses_custom_message() -> None: + form = ComparisonForm(MultiDict({"source": "same", "target": "same"})) + with pytest.raises(ValidationError, match="Must differ from Target"): + NotEqualTo("target", "Must differ from %(other_label)s")(form, form.source) + + +def test_not_equal_to_rejects_unknown_field() -> None: + form = ComparisonForm(MultiDict({"source": "source", "target": "target"})) + with pytest.raises(ValidationError, match="Invalid field name 'missing'"): + NotEqualTo("missing")(form, form.source) + + +def test_allowed_hostname_accepts_configured_hostname(app: Flask) -> None: + app.config["WEBMENTIONS_SSG_ALLOWED_HOSTNAMES"] = ["dennisfink.me"] + form = URLForm(MultiDict({"url": "https://dennisfink.me/blog/example/"})) + with app.app_context(): + AllowedHostname()(form, form.url) + + +def test_allowed_hostname_rejects_unconfigured_hostname(app: Flask) -> None: + app.config["WEBMENTIONS_SSG_ALLOWED_HOSTNAMES"] = ["dennisfink.me"] + form = URLForm(MultiDict({"url": "https://example.com/post"})) + with app.app_context(), pytest.raises(ValidationError, match="Invalid input"): + AllowedHostname()(form, form.url) + + +def test_allowed_hostname_uses_custom_message(app: Flask) -> None: + app.config["WEBMENTIONS_SSG_ALLOWED_HOSTNAMES"] = ["dennisfink.me"] + form = URLForm(MultiDict({"url": "https://example.com/post"})) + with ( + app.app_context(), + pytest.raises(ValidationError, match="Hostname is not allowed"), + ): + AllowedHostname("Hostname is not allowed")(form, form.url) + + +def test_public_url_accepts_public_url(monkeypatch: pytest.MonkeyPatch) -> None: + is_public_url = Mock(return_value=True) + monkeypatch.setattr("webmentions_ssg.forms.validators.is_public_url", is_public_url) + form = URLForm(MultiDict({"url": "https://example.com/post"})) + PublicURL()(form, form.url) + is_public_url.assert_called_once_with("https://example.com/post") + + +@pytest.mark.parametrize( + "outcome", + [ + pytest.param(False, id="non-public-address"), + pytest.param( + AddressResolutionError("Could not resolve hostname"), id="resolution-error" + ), + pytest.param(ValueError("No hostname was specified"), id="missing-hostname"), + ], +) +def test_public_url_rejects_invalid_url( + monkeypatch: pytest.MonkeyPatch, outcome: bool | Exception +) -> None: + is_public_url = Mock() + + if isinstance(outcome, Exception): + is_public_url.side_effect = outcome + else: + is_public_url.return_value = outcome + + monkeypatch.setattr("webmentions_ssg.forms.validators.is_public_url", is_public_url) + form = URLForm(MultiDict({"url": "https://example.com/post"})) + + with pytest.raises(ValidationError, match="URL must resolve to a public address"): + PublicURL()(form, form.url) + + +def test_public_url_uses_custom_message(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + "webmentions_ssg.forms.validators.is_public_url", Mock(return_value=False) + ) + form = URLForm(MultiDict({"url": "https://example.com/post"})) + + with pytest.raises(ValidationError, match="Public URL required"): + PublicURL("Public URL required")(form, form.url) diff --git a/tests/tasks/test_extension.py b/tests/tasks/test_extension.py new file mode 100644 index 0000000..5e9cdfa --- /dev/null +++ b/tests/tasks/test_extension.py @@ -0,0 +1,175 @@ +# SPDX-FileCopyrightText: 2026 Dennis Fink <me+coding@dennisfink.me> +# +# SPDX-License-Identifier: BSD-3-Clause + +from collections.abc import Callable +from typing import Any + +import huey as huey_package +import pytest +from flask import Flask, has_app_context +from huey import crontab + +from webmentions_ssg.tasks.extension import Huey + + +def test_init_without_app() -> None: + huey = Huey() + assert huey.app is None + with pytest.raises(RuntimeError, match="Huey has not been initialized"): + _ = huey.huey + + +def test_init_with_app(app: Flask) -> None: + huey = Huey(app) + assert huey.app is app + assert huey.huey.name == app.import_name + assert huey.huey.results is True + assert huey.huey.store_none is False + assert huey.huey.utc is True + assert huey.huey.immediate is True + assert app.extensions["huey"] is huey + + +def test_init_app_uses_huey_configuration(app: Flask) -> None: + app.config["HUEY"] = { + "url": "memory://", + "name": "custom-name", + "results": False, + "store_none": True, + "utc": False, + "immediate": False, + } + + huey = Huey() + huey.init_app(app) + + assert huey.huey.name == "custom-name" + assert huey.huey.results is False + assert huey.huey.store_none is True + assert huey.huey.utc is False + assert huey.huey.immediate is False + + +def test_huey_url_takes_precedence_over_huey_configuration(app: Flask) -> None: + app.config["HUEY"] = {"url": "blackhole://"} + app.config["HUEY_URL"] = "memory://" + + huey = Huey(app) + + assert type(huey.huey).__name__ == "MemoryHuey" + + +def task_decorator(huey: Huey, periodic: bool) -> Callable[[Callable[..., Any]], Any]: + if periodic: + return huey.periodic_task(crontab(minute="*")) + + return huey.task() + + +@pytest.mark.parametrize("periodic", [False, True], ids=["task", "periodic-task"]) +def test_task_runs_in_app_context(app: Flask, periodic: bool) -> None: + huey = Huey(app) + + @task_decorator(huey, periodic) + def add(a: int, b: int) -> int: + assert has_app_context() + return a + b + + assert add.call_local(2, 3) == 5 + + +@pytest.mark.parametrize("periodic", [False, True], ids=["task", "periodic-task"]) +def test_task_raises_without_app(app: Flask, periodic: bool) -> None: + huey = Huey(app) + + @task_decorator(huey, periodic) + def task() -> None: + pass + + huey.app = None + + with pytest.raises(RuntimeError, match="Flask app is not available"): + task.call_local() + + +def test_getattr_delegates_to_huey(app: Flask) -> None: + huey = Huey(app) + assert huey.immediate is True + + +@pytest.mark.parametrize( + ("url", "backend_name", "storage_kwargs"), + [ + ("redis://localhost:6379/0", "RedisHuey", {"url": "redis://localhost:6379/0"}), + ( + "rediss://localhost:6379/0", + "RedisHuey", + {"url": "rediss://localhost:6379/0"}, + ), + ( + "redis+priority://localhost:6379/0", + "PriorityRedisHuey", + {"url": "redis://localhost:6379/0"}, + ), + ( + "redis+expire://localhost:6379/0", + "RedisExpireHuey", + {"url": "redis://localhost:6379/0"}, + ), + ( + "redis+priority+expire://localhost:6379/0", + "PriorityRedisExpireHuey", + {"url": "redis://localhost:6379/0"}, + ), + ("sqlite:///var/huey.db", "SqliteHuey", {"filename": "var/huey.db"}), + ("file:///var/huey-queue", "FileHuey", {"path": "var/huey-queue"}), + ( + "postgres://user:password@localhost/database", + "PostgresHuey", + {"dsn": "postgres://user:password@localhost/database"}, + ), + ( + "postgresql://user:password@localhost/database", + "PostgresHuey", + {"dsn": "postgresql://user:password@localhost/database"}, + ), + ("memory://", "MemoryHuey", {}), + ("blackhole://", "BlackHoleHuey", {}), + ], +) +def test_backend_from_url( + monkeypatch: pytest.MonkeyPatch, + url: str, + backend_name: str, + storage_kwargs: dict[str, str], +) -> None: + backend_class = type(f"Test{backend_name}", (), {}) + + monkeypatch.setattr(huey_package, backend_name, backend_class) + + actual_class, actual_storage_kwargs = Huey.backend_from_url(url) + + assert actual_class is backend_class + assert actual_storage_kwargs == storage_kwargs + + +@pytest.mark.parametrize( + ("url", "message"), + [ + ( + "sqlite://var/huey.db", + "SQLite Huey URLs must look like sqlite:///var/huey.db", + ), + ("sqlite:///", "SQLite Huey URL must include a database path"), + ( + "file://var/huey-queue", + "File Huey URLs must look like file:///var/huey-queue", + ), + ("file:///", "File Huey URL must include a directory path"), + ("amqp://localhost", "Unsupported HUEY_URL scheme: 'amqp'"), + ], +) +def test_backend_from_url_rejects_invalid_url(url: str, message: str) -> None: + with pytest.raises(RuntimeError, match=message): + Huey.backend_from_url(url) diff --git a/tests/tasks/test_receiver.py b/tests/tasks/test_receiver.py index 17dd45a..daa271d 100644 --- a/tests/tasks/test_receiver.py +++ b/tests/tasks/test_receiver.py @@ -54,9 +54,7 @@ def create_webmention( def get_webmention_state(app: Flask, identifier: uuid.UUID) -> tuple[str, str | None]: with app.app_context(): webmention = db.session.get(ReceivedWebmention, identifier) - assert webmention is not None - return (webmention.status, webmention.failure_reason) @@ -433,9 +431,7 @@ def test_verify_webmention_marks_row_verifying_before_check( return True source_mentions_target.side_effect = verify_source - receiver.verify_webmention.call_local(identifier) - assert get_webmention_state(app, identifier) == ("verified", None) @@ -521,7 +517,6 @@ def test_ensure_public_request_accepts_public_url( is_public_url: Mock, receiver: ModuleType ) -> None: receiver.ensure_public_request(httpx.Request("GET", SOURCE_URL)) - is_public_url.assert_called_once_with(SOURCE_URL) diff --git a/tests/tasks/test_scanner.py b/tests/tasks/test_scanner.py index 332b8bf..3de5e16 100644 --- a/tests/tasks/test_scanner.py +++ b/tests/tasks/test_scanner.py @@ -4,6 +4,7 @@ from pathlib import Path from types import ModuleType +from unittest.mock import Mock from uuid import UUID import pytest @@ -20,13 +21,11 @@ SOURCE_URL = f"{BASE_URL}example/" @pytest.fixture -def scanner_module(app: Flask) -> ModuleType: - # Importing scanner registers Huey tasks. The dependency on the - # app fixture guarantees that Huey has been initialized first. +def scanner_module(app: Flask, monkeypatch: pytest.MonkeyPatch) -> ModuleType: _ = app - from webmentions_ssg.tasks import scanner + monkeypatch.setattr(scanner, "is_http_url", Mock(return_value=True)) return scanner @@ -96,11 +95,6 @@ def configure_scanner(app: Flask, root: Path, *, base_url: str | None = None) -> app.config["WEBMENTIONS_SSG_SOURCE_BASE_URL"] = base_url -# --------------------------------------------------------------------------- -# Microformats parsing -# --------------------------------------------------------------------------- - - def test_parse_entry_returns_microformats_entry(scanner_module: ModuleType) -> None: document = parse_html( f""" @@ -277,11 +271,6 @@ def test_primary_entry_rejects_nonmatching_u_url(scanner_module: ModuleType) -> scanner_module.primary_entry(document, SOURCE_URL) -# --------------------------------------------------------------------------- -# e-content -# --------------------------------------------------------------------------- - - def test_content_element_returns_e_content(scanner_module: ModuleType) -> None: document = parse_html( """ @@ -346,11 +335,6 @@ def test_content_element_rejects_multiple_e_content(scanner_module: ModuleType) scanner_module.content_element(entry) -# --------------------------------------------------------------------------- -# Canonical URLs -# --------------------------------------------------------------------------- - - def test_canonical_url_returns_absolute_url(scanner_module: ModuleType) -> None: document = parse_html( f""" @@ -393,6 +377,8 @@ def test_canonical_url_returns_none_when_missing(scanner_module: ModuleType) -> def test_canonical_url_resolves_relative_url_with_base_url( scanner_module: ModuleType, ) -> None: + scanner_module.is_http_url.side_effect = [False, True] + document = parse_html( """ <html> @@ -414,6 +400,8 @@ def test_canonical_url_resolves_relative_url_with_base_url( def test_relative_canonical_requires_base_url(scanner_module: ModuleType) -> None: + scanner_module.is_http_url.return_value = False + document = parse_html( """ <html> @@ -453,11 +441,6 @@ def test_empty_canonical_is_rejected(scanner_module: ModuleType) -> None: ) -# --------------------------------------------------------------------------- -# Target extraction -# --------------------------------------------------------------------------- - - def test_scan_source_uses_canonical_without_base_url( scanner_module: ModuleType, tmp_path: Path ) -> None: @@ -657,24 +640,6 @@ def test_scan_source_extracts_reaction_properties( assert scanned.targets == frozenset({target}) -def test_scan_source_extracts_anchor_from_e_content( - scanner_module: ModuleType, tmp_path: Path -) -> None: - path = write_post( - tmp_path, - "example", - """ - <a href="https://example.com/target"> - Target - </a> - """, - ) - - scanned = scanner_module.scan_source_file(path, root=tmp_path, base_url=None) - - assert scanned.targets == frozenset({"https://example.com/target"}) - - @pytest.mark.parametrize("tag", ("link", "area")) def test_scan_source_ignores_non_anchor_href_elements( scanner_module: ModuleType, tmp_path: Path, tag: str @@ -714,11 +679,6 @@ def test_scan_source_ignores_self_fragment( assert scanned.targets == frozenset({"https://example.com/target"}) -# --------------------------------------------------------------------------- -# Database reconciliation -# --------------------------------------------------------------------------- - - def test_scan_creates_source_and_webmentions( app: Flask, scanner_module: ModuleType, tmp_path: Path, monkeypatch: MonkeyPatch ) -> None: @@ -900,7 +860,6 @@ def test_updated_source_increments_revision( assert webmention.desired_revision == 2 assert webmention.processed_revision == 1 assert webmention.sent_revision == 1 - assert webmention.pending identifier = webmention.uuid @@ -972,7 +931,6 @@ def test_new_target_is_added_on_update( assert webmention.desired_revision == 2 assert webmention.processed_revision is None assert webmention.sent_revision is None - assert webmention.pending identifier = webmention.uuid @@ -1055,7 +1013,6 @@ def test_removed_sent_target_is_queued_again( assert webmention.desired_revision == 2 assert webmention.processed_revision == 1 assert webmention.sent_revision == 1 - assert webmention.pending assert queued == [identifier] @@ -1124,16 +1081,10 @@ def test_removed_unsent_target_is_not_queued( assert webmention.desired_revision == 2 assert webmention.processed_revision == 2 assert webmention.sent_revision is None - assert not webmention.pending assert queued == [] -# --------------------------------------------------------------------------- -# Source deletion/restoration -# --------------------------------------------------------------------------- - - def test_deleted_source_queues_previously_sent_webmention( app: Flask, scanner_module: ModuleType, tmp_path: Path, monkeypatch: MonkeyPatch ) -> None: @@ -1186,7 +1137,6 @@ def test_deleted_source_queues_previously_sent_webmention( assert webmention.desired_revision == 2 assert webmention.processed_revision == 1 assert webmention.sent_revision == 1 - assert webmention.pending assert queued == [identifier] @@ -1231,7 +1181,6 @@ def test_deleted_source_does_not_queue_unsent_webmention( assert webmention.desired_revision == 2 assert webmention.processed_revision == 2 assert webmention.sent_revision is None - assert not webmention.pending assert queued == [] @@ -1298,18 +1247,12 @@ def test_restored_source_creates_new_revision( assert webmention.active assert webmention.desired_revision == 3 assert webmention.sent_revision == 1 - assert webmention.pending identifier = webmention.uuid assert queued == [identifier] -# --------------------------------------------------------------------------- -# Queue recovery and invalid sources -# --------------------------------------------------------------------------- - - def test_pending_webmention_is_requeued_on_next_scan( app: Flask, scanner_module: ModuleType, tmp_path: Path, monkeypatch: MonkeyPatch ) -> None: @@ -1436,6 +1379,8 @@ def test_scan_sources_accepts_valid_base_url( def test_scan_sources_rejects_invalid_base_url( app: Flask, scanner_module: ModuleType, tmp_path: Path ) -> None: + scanner_module.is_http_url.return_value = False + configure_scanner(app, tmp_path, base_url="not-a-url") with pytest.raises(RuntimeError, match="absolute HTTP or HTTPS URL"): @@ -1565,6 +1510,8 @@ def test_scan_source_ignores_reaction_to_matching_hostname( def test_canonical_url_rejects_non_http_resolved_url( scanner_module: ModuleType, ) -> None: + scanner_module.is_http_url.return_value = False + document = parse_html( """ <html> @@ -1634,19 +1581,21 @@ def test_parse_entry_rejects_invalid_parser_output( scanner_module.parse_entry(element, SOURCE_URL) -@pytest.mark.parametrize( - "value", - [ - pytest.param(" ", id="empty"), - pytest.param("mailto:example@example.com", id="non-http"), - ], -) -def test_normalize_target_ignores_invalid_target( - scanner_module: ModuleType, value: str +def test_normalize_target_ignores_empty_target(scanner_module: ModuleType) -> None: + assert ( + scanner_module.normalize_target(" ", base_url=SOURCE_URL, source_url=SOURCE_URL) + is None + ) + + +def test_normalize_target_ignores_url_rejected_by_validator( + scanner_module: ModuleType, ) -> None: + scanner_module.is_http_url.return_value = False + assert ( scanner_module.normalize_target( - value, base_url=SOURCE_URL, source_url=SOURCE_URL + "https://example.com/target", base_url=SOURCE_URL, source_url=SOURCE_URL ) is None ) diff --git a/tests/tasks/test_sender.py b/tests/tasks/test_sender.py index dac56a5..6e33887 100644 --- a/tests/tasks/test_sender.py +++ b/tests/tasks/test_sender.py @@ -90,16 +90,60 @@ def mock_sender_requests( @pytest.fixture def sender_module(app: Flask) -> ModuleType: - # Importing sender registers Huey tasks. The dependency on the - # app fixture guarantees that Huey has been initialized first. _ = app - from webmentions_ssg.tasks import sender return sender @pytest.mark.parametrize( + ("status", "expected"), + [ + pytest.param(200, False, id="success"), + pytest.param(400, False, id="client-error"), + pytest.param(408, True, id="request-timeout"), + pytest.param(425, True, id="too-early"), + pytest.param(429, True, id="too-many-requests"), + pytest.param(500, True, id="server-error-lower-bound"), + pytest.param(599, True, id="server-error-upper-bound"), + pytest.param(600, False, id="outside-server-error-range"), + ], +) +def test_temporary_http_status( + sender_module: ModuleType, status: int, expected: bool +) -> None: + assert sender_module.temporary_http_status(status) is expected + + +@pytest.mark.parametrize( + ("desired_revision", "processed_revision", "revision", "expected"), + [ + pytest.param(2, None, 2, True, id="unprocessed-current"), + pytest.param(2, 1, 2, True, id="older-processed-current"), + pytest.param(2, 2, 2, False, id="already-processed"), + pytest.param(2, 3, 2, False, id="newer-processed"), + pytest.param(3, 1, 2, False, id="superseded"), + ], +) +def test_attempt_is_current( + app: Flask, + sender_module: ModuleType, + desired_revision: int, + processed_revision: int | None, + revision: int, + expected: bool, +) -> None: + identifier = create_sent_webmention( + app, desired_revision=desired_revision, processed_revision=processed_revision + ) + + with app.app_context(): + webmention = db.session.get(SentWebmention, identifier) + assert webmention is not None + assert sender_module.attempt_is_current(webmention, revision) is expected + + +@pytest.mark.parametrize( ("value", "expected"), ( ( @@ -119,8 +163,10 @@ def sender_module(app: Flask) -> ModuleType: ("https://example.com/webmention", {"webmention", "alternate"}), ), ( - "<https://example.com/webmention>; " - 'type="text/html"; rel="webmention"; title="Endpoint"', + ( + '<https://example.com/webmention>; type="text/html"; ' + 'rel="webmention"; title="Endpoint"' + ), ("https://example.com/webmention", {"webmention"}), ), ( @@ -159,7 +205,10 @@ def test_parse_link_value_rejects_malformed_value( ( ( "Link", - "<https://example.com/test/2/webmention?head=true>; rel=webmention", + ( + "<https://example.com/test/2/webmention?head=true>;" + "rel=webmention" + ), ), ), "", @@ -223,7 +272,10 @@ def test_parse_link_value_rejects_malformed_value( ( ( "LinK", - "<https://example.com/test/7/webmention?head=true>; rel=webmention", + ( + "<https://example.com/test/7/webmention?head=true>; " + "rel=webmention" + ), ), ), "", @@ -235,7 +287,10 @@ def test_parse_link_value_rejects_malformed_value( ( ( "Link", - '<https://example.com/test/8/webmention?head=true>; rel="webmention"', + ( + "<https://example.com/test/8/webmention?head=true>; " + 'rel="webmention"' + ), ), ), "", @@ -259,8 +314,10 @@ def test_parse_link_value_rejects_malformed_value( ( ( "Link", - "<https://example.com/test/10/webmention?head=true>; " - 'rel="webmention somethingelse"', + ( + "<https://example.com/test/10/webmention?head=true>; " + 'rel="webmention somethingelse"' + ), ), ), "", @@ -407,8 +464,10 @@ def test_parse_link_value_rejects_malformed_value( ("Link", '<https://example.com/test/18/webmention/error>; rel="other"'), ( "Link", - "<https://example.com/test/18/webmention?head=true>; " - 'rel="webmention"', + ( + "<https://example.com/test/18/webmention?head=true>; " + 'rel="webmention"' + ), ), ), "", @@ -420,10 +479,12 @@ def test_parse_link_value_rejects_malformed_value( ( ( "Link", - "<https://example.com/test/19/webmention/error>; " - 'rel="other", ' - "<https://example.com/test/19/webmention?head=true>; " - 'rel="webmention"', + ( + "<https://example.com/test/19/webmention/error>; " + 'rel="other", ' + "<https://example.com/test/19/webmention?head=true>; " + 'rel="webmention"' + ), ), ), "", @@ -625,7 +686,6 @@ def test_send_webmention_persists_success( assert webmention.status_url == STATUS_URL assert webmention.last_attempted_at is not None assert webmention.last_sent_at is not None - assert not webmention.pending def test_send_webmention_marks_unsupported_target_processed( @@ -663,7 +723,6 @@ def test_send_webmention_marks_unsupported_target_processed( assert webmention.response_status is None assert webmention.status_url is None assert webmention.last_attempted_at is not None - assert not webmention.pending def test_send_webmention_persists_permanent_failure( @@ -693,7 +752,6 @@ def test_send_webmention_persists_permanent_failure( assert webmention.endpoint == ENDPOINT_URL assert webmention.response_status == 400 assert webmention.status_url is None - assert not webmention.pending def test_send_webmention_persists_temporary_failure_and_reraises( @@ -727,7 +785,6 @@ def test_send_webmention_persists_temporary_failure_and_reraises( assert webmention.endpoint == ENDPOINT_URL assert webmention.response_status == 503 assert webmention.status_url is None - assert webmention.pending def test_send_webmention_ignores_processed_revision( @@ -750,14 +807,17 @@ def test_send_webmention_ignores_processed_revision( post.assert_not_called() -def test_send_webmention_ignores_result_for_superseded_revision( - app: Flask, sender_module: ModuleType, monkeypatch: pytest.MonkeyPatch +@pytest.mark.parametrize( + "outcome", ["success", "unsupported", "temporary-failure", "permanent-failure"] +) +def test_send_webmention_ignores_superseded_attempt( + app: Flask, sender_module: ModuleType, monkeypatch: pytest.MonkeyPatch, outcome: str ) -> None: identifier = create_sent_webmention(app) - _, _, _, post = mock_sender_requests(sender_module, monkeypatch) + _, _, discover, post = mock_sender_requests(sender_module, monkeypatch) - def supersede_revision(*args, **kwargs) -> tuple[int, None]: + def supersede_revision() -> None: webmention = db.session.get(SentWebmention, identifier) assert webmention is not None @@ -765,9 +825,40 @@ def test_send_webmention_ignores_result_for_superseded_revision( webmention.desired_revision = 2 db.session.commit() - return (202, None) + match outcome: + case "success": + + def post_success(*args, **kwargs) -> tuple[int, None]: + supersede_revision() + return (202, None) + + post.side_effect = post_success + + case "unsupported": + + def discover_unsupported(*args, **kwargs) -> None: + supersede_revision() + + discover.side_effect = discover_unsupported - post.side_effect = supersede_revision + case "temporary-failure": + + def post_temporary_failure(*args, **kwargs) -> None: + supersede_revision() + raise sender_module.TemporarySenderError("Temporary failure") + + post.side_effect = post_temporary_failure + + case "permanent-failure": + + def post_permanent_failure(*args, **kwargs) -> tuple[int, None]: + supersede_revision() + return (400, None) + + post.side_effect = post_permanent_failure + + case _: + raise AssertionError(f"Unexpected outcome: {outcome}") sender_module.send_webmention.call_local(identifier) @@ -779,17 +870,14 @@ def test_send_webmention_ignores_result_for_superseded_revision( assert webmention.processed_revision is None assert webmention.sent_revision is None assert webmention.status is None - assert webmention.pending def test_send_webmention_ignores_unknown_identifier( sender_module: ModuleType, caplog: pytest.LogCaptureFixture ) -> None: identifier = uuid.uuid7() - with caplog.at_level(logging.WARNING): sender_module.send_webmention.call_local(identifier) - assert f"Cannot send unknown SentWebmention {identifier}" in caplog.text @@ -798,9 +886,7 @@ def test_ensure_public_request_accepts_public_url( ) -> None: is_public_url = Mock(return_value=True) monkeypatch.setattr(sender_module, "is_public_url", is_public_url) - sender_module.ensure_public_request(httpx.Request("GET", TARGET_URL)) - is_public_url.assert_called_once_with(TARGET_URL) @@ -809,13 +895,11 @@ def test_ensure_public_request_rejects_non_public_address( ) -> None: is_public_url = Mock(return_value=False) monkeypatch.setattr(sender_module, "is_public_url", is_public_url) - with pytest.raises( sender_module.PermanentSenderError, match="Request resolves to a non-public address", ): sender_module.ensure_public_request(httpx.Request("GET", TARGET_URL)) - is_public_url.assert_called_once_with(TARGET_URL) @@ -826,23 +910,25 @@ def test_ensure_public_request_maps_dns_failure_to_temporary_error( side_effect=AddressResolutionError("Could not resolve hostname") ) monkeypatch.setattr(sender_module, "is_public_url", is_public_url) - with pytest.raises( sender_module.TemporarySenderError, match="Request hostname could not be resolved", ): sender_module.ensure_public_request(httpx.Request("GET", TARGET_URL)) - is_public_url.assert_called_once_with(TARGET_URL) -def test_resolve_endpoint_rejects_non_http_url(sender_module: ModuleType) -> None: +def test_resolve_endpoint_rejects_invalid_url( + sender_module: ModuleType, monkeypatch: pytest.MonkeyPatch +) -> None: + is_http_url = Mock(return_value=False) + monkeypatch.setattr(sender_module, "is_http_url", is_http_url) response = httpx.Response(200, request=httpx.Request("GET", TARGET_URL)) - with pytest.raises( sender_module.PermanentSenderError, match="Invalid Webmention endpoint" ): - sender_module.resolve_endpoint(response, "mailto:example@example.com") + sender_module.resolve_endpoint(response, "/webmention") + is_http_url.assert_called_once_with(ENDPOINT_URL) def test_endpoint_from_headers_ignores_malformed_link( @@ -853,7 +939,6 @@ def test_endpoint_from_headers_ignores_malformed_link( headers=[("Link", "not-a-link"), ("Link", "</webmention>; rel=webmention")], request=httpx.Request("GET", TARGET_URL), ) - assert sender_module.endpoint_from_headers(response) == ENDPOINT_URL @@ -1065,101 +1150,3 @@ def test_post_webmention( assert form["source"] == SOURCE_URL assert form["target"] == TARGET_URL - - -def test_send_webmention_ignores_unsupported_result_for_superseded_revision( - app: Flask, sender_module: ModuleType, monkeypatch: pytest.MonkeyPatch -) -> None: - identifier = create_sent_webmention(app) - - _, _, discover, post = mock_sender_requests( - sender_module, monkeypatch, endpoint=None - ) - - def supersede_revision(*args, **kwargs) -> None: - webmention = db.session.get(SentWebmention, identifier) - - assert webmention is not None - - webmention.desired_revision = 2 - db.session.commit() - - discover.side_effect = supersede_revision - - sender_module.send_webmention.call_local(identifier) - - post.assert_not_called() - - with app.app_context(): - webmention = db.session.get(SentWebmention, identifier) - - assert webmention is not None - assert webmention.desired_revision == 2 - assert webmention.processed_revision is None - assert webmention.sent_revision is None - assert webmention.status is None - assert webmention.pending - - -def test_send_webmention_ignores_temporary_failure_for_superseded_revision( - app: Flask, sender_module: ModuleType, monkeypatch: pytest.MonkeyPatch -) -> None: - identifier = create_sent_webmention(app) - - _, _, _, post = mock_sender_requests(sender_module, monkeypatch) - - def supersede_revision(*args, **kwargs) -> tuple[int, None]: - webmention = db.session.get(SentWebmention, identifier) - - assert webmention is not None - - webmention.desired_revision = 2 - db.session.commit() - - raise sender_module.TemporarySenderError("Temporary failure") - - post.side_effect = supersede_revision - - sender_module.send_webmention.call_local(identifier) - - with app.app_context(): - webmention = db.session.get(SentWebmention, identifier) - - assert webmention is not None - assert webmention.desired_revision == 2 - assert webmention.processed_revision is None - assert webmention.sent_revision is None - assert webmention.status is None - assert webmention.pending - - -def test_send_webmention_ignores_permanent_failure_for_superseded_revision( - app: Flask, sender_module: ModuleType, monkeypatch: pytest.MonkeyPatch -) -> None: - identifier = create_sent_webmention(app) - - _, _, _, post = mock_sender_requests(sender_module, monkeypatch) - - def supersede_revision(*args, **kwargs) -> tuple[int, None]: - webmention = db.session.get(SentWebmention, identifier) - - assert webmention is not None - - webmention.desired_revision = 2 - db.session.commit() - - return (400, None) - - post.side_effect = supersede_revision - - sender_module.send_webmention.call_local(identifier) - - with app.app_context(): - webmention = db.session.get(SentWebmention, identifier) - - assert webmention is not None - assert webmention.desired_revision == 2 - assert webmention.processed_revision is None - assert webmention.sent_revision is None - assert webmention.status is None - assert webmention.pending diff --git a/tests/test_models.py b/tests/test_models.py new file mode 100644 index 0000000..2fa2cb8 --- /dev/null +++ b/tests/test_models.py @@ -0,0 +1,148 @@ +# SPDX-FileCopyrightText: 2026 Dennis Fink <me+coding@dennisfink.me> +# +# SPDX-License-Identifier: BSD-3-Clause + +import uuid +from collections.abc import Callable +from datetime import UTC, datetime + +import pytest + +from webmentions_ssg.models import ( + ReceivedWebmention, + SentWebmention, + Source, + User, + uuid7_to_datetime, +) + + +def test_user_repr() -> None: + user = User(username="admin", password="secret") + assert repr(user) == "<User admin>" + + +def test_user_password_is_write_only() -> None: + user = User(username="admin", password="secret") + with pytest.raises(AttributeError, match="Password is write-only"): + _ = user.password + + +def test_user_password_is_hashed_and_can_be_checked() -> None: + user = User(username="admin", password="secret") + assert user.password_hash != "secret" + assert user.check_password("secret") + assert not user.check_password("wrong") + + +@pytest.mark.parametrize( + ("status", "expected"), + [ + ("received", False), + ("verifying", False), + ("verified", True), + ("failed", False), + ("deleted", False), + ], +) +def test_received_webmention_verified(status: str, expected: bool) -> None: + webmention = ReceivedWebmention( + uuid=uuid.uuid7(), + source="https://source.example/post", + target="https://dennisfink.me/blog/example/", + status=status, + ) + + assert webmention.verified is expected + + +@pytest.mark.parametrize( + ("desired_revision", "processed_revision", "expected"), + [(1, None, True), (1, 0, True), (2, 1, True), (1, 1, False), (1, 2, False)], +) +def test_sent_webmention_pending( + desired_revision: int, processed_revision: int | None, expected: bool +) -> None: + webmention = SentWebmention( + target="https://example.com/post", + desired_revision=desired_revision, + processed_revision=processed_revision, + ) + assert webmention.pending is expected + + +def test_uuid7_to_datetime() -> None: + identifier = uuid.uuid7() + timestamp = uuid7_to_datetime(identifier) + + assert timestamp.tzinfo is UTC + assert timestamp.timestamp() == pytest.approx(identifier.time / 1000) + + +@pytest.mark.parametrize( + "factory", + [ + pytest.param( + lambda identifier: ReceivedWebmention( + uuid=identifier, + source="https://source.example/post", + target="https://dennisfink.me/blog/example/", + ), + id="received-webmention", + ), + pytest.param( + lambda identifier: Source( + uuid=identifier, + path="example/index.html", + url="https://dennisfink.me/blog/example/", + content_hash="0" * 32, + revision=1, + last_seen_at=datetime.now(UTC), + revised_at=datetime.now(UTC), + ), + id="source", + ), + pytest.param( + lambda identifier: SentWebmention( + uuid=identifier, target="https://example.com/post", desired_revision=1 + ), + id="sent-webmention", + ), + ], +) +def test_created_at_comes_from_uuid7_timestamp( + factory: Callable[[uuid.UUID], ReceivedWebmention | Source | SentWebmention], +) -> None: + identifier = uuid.uuid7() + model = factory(identifier) + assert model.created_at == uuid7_to_datetime(identifier) + + +def test_received_webmention_repr() -> None: + webmention = ReceivedWebmention( + uuid=uuid.uuid7(), + source="https://source.example/post", + target="https://dennisfink.me/blog/example/", + ) + + assert repr(webmention) == ( + "<ReceivedWebmention 'https://source.example/post' -> " + "'https://dennisfink.me/blog/example/'>" + ) + + +def test_source_repr() -> None: + source = Source( + path="example/index.html", + url="https://dennisfink.me/blog/example/", + content_hash="0" * 32, + revision=1, + last_seen_at=datetime.now(UTC), + revised_at=datetime.now(UTC), + ) + assert repr(source) == "<Source 'https://dennisfink.me/blog/example/'>" + + +def test_sent_webmention_repr() -> None: + webmention = SentWebmention(target="https://example.com/post", desired_revision=1) + assert repr(webmention) == "<SentWebmention 'https://example.com/post'>" diff --git a/tests/test_url_security.py b/tests/test_url_security.py index db513f9..dee38dd 100644 --- a/tests/test_url_security.py +++ b/tests/test_url_security.py @@ -14,6 +14,27 @@ from webmentions_ssg.url_security import ( ) +@pytest.mark.parametrize( + ("url", "expected"), + [ + pytest.param("http://example.com/", True, id="http"), + pytest.param("https://example.com/path", True, id="https"), + pytest.param("HTTP://EXAMPLE.COM/", True, id="scheme-case-insensitive"), + pytest.param("https://example.com:8443/path", True, id="port"), + pytest.param("http://[2001:db8::1]/", True, id="ipv6"), + pytest.param("ftp://example.com/", False, id="ftp"), + pytest.param("mailto:example@example.com", False, id="mailto"), + pytest.param("/relative/url", False, id="relative"), + pytest.param("//example.com/path", False, id="scheme-relative"), + pytest.param("https:///missing-host", False, id="missing-host"), + pytest.param("", False, id="empty"), + pytest.param("http://[::1", False, id="malformed"), + ], +) +def test_is_http_url(url: str, expected: bool) -> None: + assert is_http_url(url) is expected + + @patch("webmentions_ssg.url_security.dns.resolver.resolve_name") @pytest.mark.parametrize( ("url", "resolved_addresses"), @@ -34,7 +55,6 @@ def test_is_public_url_returns_false_for_non_public_addresses( resolve_name: Mock, url: str, resolved_addresses: list[str] ) -> None: resolve_name.return_value.addresses.return_value = resolved_addresses - assert not is_public_url(url) @@ -53,7 +73,6 @@ def test_is_public_url_returns_true_for_public_addresses( resolve_name: Mock, url: str, resolved_addresses: list[str] ) -> None: resolve_name.return_value.addresses.return_value = resolved_addresses - assert is_public_url(url) @@ -64,7 +83,6 @@ def test_is_public_url_returns_true_for_public_addresses( def test_is_public_url_raises_for_resolution_failure(resolve_name: Mock) -> None: with pytest.raises(AddressResolutionError, match="Could not resolve hostname"): is_public_url("https://nonexistent.example/") - resolve_name.assert_called_once_with("nonexistent.example") @@ -72,9 +90,4 @@ def test_is_public_url_raises_for_resolution_failure(resolve_name: Mock) -> None def test_is_public_url_raises_when_url_has_no_hostname(resolve_name: Mock) -> None: with pytest.raises(ValueError, match="No hostname was specified"): is_public_url("/relative/url") - resolve_name.assert_not_called() - - -def test_is_http_url_rejects_malformed_url() -> None: - assert not is_http_url("http://[::1") diff --git a/tests/test_views.py b/tests/test_views.py index 0de8b9b..844da2e 100644 --- a/tests/test_views.py +++ b/tests/test_views.py @@ -4,7 +4,8 @@ import uuid from datetime import UTC, datetime -from unittest.mock import Mock, call, patch +from types import ModuleType +from unittest.mock import Mock, call import pytest import sqlalchemy as sa @@ -20,6 +21,36 @@ from webmentions_ssg.models import ( User, ) +SOURCE_URL = "https://source.example/post" +TARGET_URL = "https://dennisfink.me/blog/example/" + + +@pytest.fixture +def views_module(app: Flask) -> ModuleType: + _ = app + from webmentions_ssg import views + + return views + + +def mock_endpoint_form( + views_module: ModuleType, + monkeypatch: pytest.MonkeyPatch, + *, + valid: bool, + source: str = SOURCE_URL, + target: str = TARGET_URL, + errors: dict[str, list[str]] | None = None, +) -> Mock: + form = Mock() + form.validate_on_submit.return_value = valid + form.source.data = source + form.target.data = target + form.errors = errors or {} + + monkeypatch.setattr(views_module.forms, "EndpointForm", Mock(return_value=form)) + return form + def log_in(app: Flask, client: FlaskClient) -> None: with app.app_context(): @@ -82,46 +113,42 @@ def create_sent_source(app: Flask) -> uuid.UUID: def test_endpoint_only_accepts_post(client: FlaskClient) -> None: response = client.get("/endpoint") - assert response.status_code == 405 -@patch("webmentions_ssg.forms.validators.is_public_url", new=lambda url: True) -def test_endpoint_returns_form_errors(client: FlaskClient) -> None: - response = client.post( - "/endpoint", - data={ - "source": "https://source.example/post", - "target": "https://example.com/post", - }, - ) - - assert response.status_code == 400 +def test_endpoint_returns_form_errors( + client: FlaskClient, views_module: ModuleType, monkeypatch: pytest.MonkeyPatch +) -> None: + errors = {"source": ["Invalid source"]} + mock_endpoint_form(views_module, monkeypatch, valid=False, errors=errors) - errors = response.get_json() + response = client.post("/endpoint") - assert errors is not None - assert "target" in errors + assert response.status_code == 400 + assert response.get_json() == errors -@patch("webmentions_ssg.forms.validators.is_public_url", new=lambda url: True) -@patch("webmentions_ssg.views.verify_webmention") def test_endpoint_creates_webmention( - verify_webmention: Mock, app: Flask, client: FlaskClient + app: Flask, + client: FlaskClient, + views_module: ModuleType, + monkeypatch: pytest.MonkeyPatch, ) -> None: - source = "https://source.example/post" - target = "https://dennisfink.me/blog/example/" + form = mock_endpoint_form(views_module, monkeypatch, valid=True) + verify_webmention = Mock() + monkeypatch.setattr(views_module, "verify_webmention", verify_webmention) - response = client.post("/endpoint", data={"source": source, "target": target}) + response = client.post("/endpoint") assert response.status_code == 201 + form.validate_on_submit.assert_called_once_with() with app.app_context(): webmention = db.session.scalar(sa.select(ReceivedWebmention)) assert webmention is not None - assert webmention.source == source - assert webmention.target == target + assert webmention.source == SOURCE_URL + assert webmention.target == TARGET_URL assert webmention.status == "received" assert webmention.failure_reason is None @@ -131,18 +158,18 @@ def test_endpoint_creates_webmention( verify_webmention.assert_called_once_with(identifier) -@patch("webmentions_ssg.forms.validators.is_public_url", new=lambda url: True) -@patch("webmentions_ssg.views.verify_webmention") def test_endpoint_is_idempotent( - verify_webmention: Mock, app: Flask, client: FlaskClient + app: Flask, + client: FlaskClient, + views_module: ModuleType, + monkeypatch: pytest.MonkeyPatch, ) -> None: - data = { - "source": "https://source.example/post", - "target": "https://dennisfink.me/blog/example/", - } + mock_endpoint_form(views_module, monkeypatch, valid=True) + verify_webmention = Mock() + monkeypatch.setattr(views_module, "verify_webmention", verify_webmention) - first_response = client.post("/endpoint", data=data) - second_response = client.post("/endpoint", data=data) + first_response = client.post("/endpoint") + second_response = client.post("/endpoint") assert first_response.status_code == 201 assert second_response.status_code == 201 @@ -155,8 +182,8 @@ def test_endpoint_is_idempotent( webmention = webmentions[0] - assert webmention.source == data["source"] - assert webmention.target == data["target"] + assert webmention.source == SOURCE_URL + assert webmention.target == TARGET_URL assert webmention.status == "received" assert webmention.failure_reason is None @@ -165,19 +192,21 @@ def test_endpoint_is_idempotent( assert verify_webmention.call_args_list == [call(identifier), call(identifier)] -@patch("webmentions_ssg.forms.validators.is_public_url", new=lambda url: True) -@patch("webmentions_ssg.views.verify_webmention") def test_resending_resets_failure_state( - verify_webmention: Mock, app: Flask, client: FlaskClient + app: Flask, + client: FlaskClient, + views_module: ModuleType, + monkeypatch: pytest.MonkeyPatch, ) -> None: - source = "https://source.example/post" - target = "https://dennisfink.me/blog/example/" + mock_endpoint_form(views_module, monkeypatch, valid=True) + verify_webmention = Mock() + monkeypatch.setattr(views_module, "verify_webmention", verify_webmention) with app.app_context(): existing = ReceivedWebmention( uuid=uuid.uuid7(), - source=source, - target=target, + source=SOURCE_URL, + target=TARGET_URL, status="failed", failure_reason="Previous failure", ) @@ -187,7 +216,7 @@ def test_resending_resets_failure_state( identifier = existing.uuid - response = client.post("/endpoint", data={"source": source, "target": target}) + response = client.post("/endpoint") assert response.status_code == 201 @@ -252,7 +281,5 @@ def test_sent_source_returns_404_for_unknown_source( app: Flask, client: FlaskClient ) -> None: log_in(app, client) - response = client.get(f"/sent/{uuid.uuid7()}") - assert response.status_code == 404 diff --git a/webmentions_ssg/tasks/extension.py b/webmentions_ssg/tasks/extension.py index 95f190f..25f8e47 100644 --- a/webmentions_ssg/tasks/extension.py +++ b/webmentions_ssg/tasks/extension.py @@ -2,22 +2,50 @@ # # SPDX-License-Identifier: BSD-3-Clause +from collections.abc import Callable from functools import wraps -from typing import Any, Callable +from typing import Any, ParamSpec, TypeVar from urllib.parse import urlsplit, urlunsplit from flask import Flask +from huey import Huey as BaseHuey +from huey.api import TaskWrapper + +P = ParamSpec("P") +R = TypeVar("R") class Huey: - def __init__(self, app: Flask | None = None): + """ + Provide Flask integration for a Huey instance. + + The extension initializes a Huey backend from the Flask configuration and + wraps tasks so that they execute within an application context. + """ + + def __init__(self, app: Flask | None = None) -> None: + """ + Initialize the Huey extension. + + :param app: Flask application to initialize immediately, if provided. + """ self.app: Flask | None = None - self._huey = None + self._huey: BaseHuey | None = None if app is not None: self.init_app(app) - def init_app(self, app: Flask): + def init_app(self, app: Flask) -> None: + """ + Initialize Huey for a Flask application. + + The backend and its storage options are derived from ``HUEY_URL`` and the + resulting extension is registered with the application. + + :param app: Flask application to initialize. + :raises RuntimeError: If the configured Huey backend URL is invalid or uses + an unsupported scheme. + """ config: dict[str, Any] = { "name": app.import_name, "results": True, @@ -36,7 +64,13 @@ class Huey: app.extensions["huey"] = self @property - def huey(self): + def huey(self) -> BaseHuey: + """ + Return the initialized Huey instance. + + :return: Configured Huey backend instance. + :raises RuntimeError: If the extension has not been initialized. + """ if self._huey is None: raise RuntimeError( "Huey has not been initialized. " @@ -44,10 +78,20 @@ class Huey: ) return self._huey - def task(self, *task_args: Any, **task_kwargs: Any): - def decorator(func: Callable): + def task( + self, *task_args: Any, **task_kwargs: Any + ) -> Callable[[Callable[P, object]], TaskWrapper]: + """ + Create a Huey task that runs within the Flask application context. + + :param task_args: Positional arguments forwarded to Huey's task decorator. + :param task_kwargs: Keyword arguments forwarded to Huey's task decorator. + :return: Decorator that registers the wrapped function as a Huey task. + """ + + def decorator(func: Callable[P, R]) -> TaskWrapper: @wraps(func) - def wrapper(*args: Any, **kwargs: Any): + def wrapper(*args: P.args, **kwargs: P.kwargs) -> R: if self.app is None: raise RuntimeError("Flask app is not available.") @@ -58,10 +102,23 @@ class Huey: return decorator - def periodic_task(self, *task_args: Any, **task_kwargs: Any): - def decorator(func: Callable): + def periodic_task( + self, *task_args: Any, **task_kwargs: Any + ) -> Callable[[Callable[P, object]], TaskWrapper]: + """ + Create a periodic Huey task that runs within the Flask application context. + + :param task_args: Positional arguments forwarded to Huey's periodic task + decorator. + :param task_kwargs: Keyword arguments forwarded to Huey's periodic task + decorator. + :return: Decorator that registers the wrapped function as a periodic Huey + task. + """ + + def decorator(func: Callable[P, R]) -> TaskWrapper: @wraps(func) - def wrapper(*args: Any, **kwargs: Any): + def wrapper(*args: P.args, **kwargs: P.kwargs) -> R: if self.app is None: raise RuntimeError("Flask app is not available.") @@ -72,23 +129,30 @@ class Huey: return decorator - def __getattr__(self, name: str): + def __getattr__(self, name: str) -> Any: """ - Forward unknown attributes to the real Huey instance. + Forward an unknown attribute to the underlying Huey instance. - This lets you still use things like: - huey.enqueue(...) - huey.scheduled() - huey.pending() + :param name: Name of the attribute to retrieve. + :return: Attribute from the initialized Huey instance. + :raises RuntimeError: If the extension has not been initialized. """ return getattr(self.huey, name) @staticmethod - def backend_from_url(url: str) -> tuple[Any, dict[str, str]]: + def backend_from_url(url: str) -> tuple[type[BaseHuey], dict[str, Any]]: + """ + Determine the Huey backend and storage options from a URL. + + :param url: Huey backend URL. + :return: Huey backend class and keyword arguments for its storage backend. + :raises RuntimeError: If the URL is malformed for the selected backend or + uses an unsupported scheme. + """ parsed = urlsplit(url) scheme = parsed.scheme.lower() - if scheme.startswith("redis") or scheme.startswith("rediss"): + if scheme.startswith(("redis", "rediss")): fixed_url = urlunsplit(parsed._replace(scheme=scheme.split("+", 1)[0])) if scheme.endswith("priority+expire"): diff --git a/webmentions_ssg/tasks/scanner.py b/webmentions_ssg/tasks/scanner.py index 0cf038a..17898ba 100644 --- a/webmentions_ssg/tasks/scanner.py +++ b/webmentions_ssg/tasks/scanner.py @@ -488,7 +488,6 @@ def scan_sources() -> None: if webmention.sent_revision is None: webmention.processed_revision = source.revision - # Persist the desired state before queueing any work. db.session.commit() pending_count = 0 |
